Compare commits
42
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
81d0cb07e2 | ||
|
|
1ea2517e22 | ||
|
|
0c63528c59 | ||
|
|
b063f8ca41 | ||
|
|
e7fff0173a | ||
|
|
d82abc271e | ||
|
|
970409962f | ||
|
|
055586703d | ||
|
|
5d89f86675 | ||
|
|
19a51a1fe6 | ||
|
|
d3232cea5a | ||
|
|
0c90c8c24d | ||
|
|
4c08ffce49 | ||
|
|
af4a77553c | ||
|
|
c096fda1eb | ||
|
|
8f47e85be0 | ||
|
|
90d3bd19eb | ||
|
|
afb4f7d3c5 | ||
|
|
02e1143f22 | ||
|
|
e2f4d1a7b5 | ||
|
|
f037351146 | ||
|
|
1ee11e08dc | ||
|
|
595f0ea60e | ||
|
|
d921832cd2 | ||
|
|
629697629a | ||
|
|
a25313beec | ||
|
|
dbde64385b | ||
|
|
9d909f5f04 | ||
|
|
76b0550c15 | ||
|
|
384c1e9493 | ||
|
|
b1dbcc93f6 | ||
|
|
b93833772e | ||
|
|
30b523edd6 | ||
|
|
6aab7f3832 | ||
|
|
9cd53fe5f8 | ||
|
|
6a32cf3a5e | ||
|
|
98be9b3da2 | ||
|
|
c53e85b767 | ||
|
|
40a8bd2d3b | ||
|
|
31aa115611 | ||
|
|
98ac10a528 | ||
|
|
a5a6d171e5 |
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
@@ -148,8 +148,7 @@ For each metric with at least two non-null source values:
|
||||
4. Stop if any source record regresses against the batch median by more than
|
||||
`max_intra_batch_regression`.
|
||||
|
||||
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
|
||||
otherwise `0.05`. Print a table with per-source values, batch median, and
|
||||
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
|
||||
worst intra-batch regression.
|
||||
|
||||
This check prevents uploading a mixed batch where one JSON is materially
|
||||
@@ -183,7 +182,7 @@ present, that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records under /tmp
|
||||
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
@@ -192,7 +191,7 @@ The sync command pattern is:
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
@@ -200,8 +199,8 @@ Then back up only the sanitized model directory under `/tmp`:
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
|
||||
from hf_store import sanitize
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.performance.hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
@@ -235,7 +234,7 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from hf_store import load_records_for_model
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
@@ -372,7 +371,7 @@ prepared records plus backup on disk.
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```python
|
||||
from hf_store import upload_record
|
||||
from fastvideo.performance.hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
@@ -460,7 +459,7 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
+28
-1
@@ -114,6 +114,17 @@ steps:
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Extraction Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
@@ -371,6 +382,21 @@ steps:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "scripts/lora_extraction/**"
|
||||
- "fastvideo/tests/lora_extraction/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/training/training_utils.py"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Extraction Tests"
|
||||
env:
|
||||
- TEST_TYPE=lora_extraction
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
@@ -410,7 +436,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
@@ -455,6 +481,7 @@ steps:
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/worker/**"
|
||||
- "fastvideo/entrypoints/**"
|
||||
- "fastvideo/performance/**"
|
||||
- "fastvideo/tests/performance/**"
|
||||
- ".buildkite/performance-benchmarks/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -80,6 +80,23 @@ MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUI
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
is_truthy() {
|
||||
case "${1:-}" in
|
||||
1|true|TRUE|yes|YES|on|ON) return 0 ;;
|
||||
*) return 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
ssim_bootstrap_args() {
|
||||
local title="${PR_TITLE:-}"
|
||||
local message="${BUILDKITE_MESSAGE:-}"
|
||||
if is_truthy "${FASTVIDEO_SSIM_BOOTSTRAP_MODE:-}" \
|
||||
|| [[ "$title" == *"[new-model]"* ]] \
|
||||
|| [[ "$message" == *"[new-model]"* ]]; then
|
||||
printf ' --bootstrap-mode'
|
||||
fi
|
||||
}
|
||||
|
||||
upload_performance_artifacts() {
|
||||
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
|
||||
LOCAL_DIR="downloaded_reports"
|
||||
@@ -172,7 +189,12 @@ case "$TEST_TYPE" in
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_SSIM_TEST_FILE::run_ssim_tests"
|
||||
SSIM_BOOTSTRAP_ARGS=$(ssim_bootstrap_args)
|
||||
if [ -n "$SSIM_BOOTSTRAP_ARGS" ]; then
|
||||
log "SSIM bootstrap mode enabled for new-model reference draft generation"
|
||||
fi
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run "
|
||||
MODAL_COMMAND+="$MODAL_SSIM_TEST_FILE::run_ssim_tests$SSIM_BOOTSTRAP_ARGS"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
|
||||
Executable
+133
@@ -0,0 +1,133 @@
|
||||
#!/usr/bin/env bash
|
||||
# Gate the expensive Buildkite full suite on the cheap GitHub checks.
|
||||
#
|
||||
# Polls the workflow runs for the PR head commit and only exits 0 once the
|
||||
# watched cheap workflows (pre-commit, docs build) have succeeded, so the
|
||||
# 'ready' label cannot burn ~20 GPU lanes on a head that a cheap check has
|
||||
# already doomed.
|
||||
#
|
||||
# Semantics:
|
||||
# - watched run completed with a bad conclusion -> exit 1 (fail CLOSED:
|
||||
# no full suite; the next push re-arms via the 'synchronize' trigger)
|
||||
# - watched run cancelled -> still pending: the docs
|
||||
# workflow's repo-global 'pages' concurrency group cancels runs superseded
|
||||
# by unrelated pushes, so 'cancelled' is not a verdict on this PR
|
||||
# - watched runs pending -> poll until done
|
||||
# - docs run absent -> not applicable after a
|
||||
# short grace period ('Deploy Documentation' is path-filtered on PRs)
|
||||
# - pre-commit run absent -> keep polling: pre-commit
|
||||
# is never path-filtered, so its absence is always anomalous
|
||||
# - 'ready' label removed while waiting -> exit 1 (fail CLOSED:
|
||||
# un-labeling is a deliberate maintainer action)
|
||||
# - GitHub API unreachable or timeout -> exit 0 (fail OPEN,
|
||||
# loud warning: never brick CI on a GitHub outage)
|
||||
#
|
||||
# Required env: PR_SHA (PR head commit), PR_NUMBER, GITHUB_REPOSITORY, GH_TOKEN.
|
||||
set -euo pipefail
|
||||
|
||||
: "${PR_SHA:?PR_SHA (PR head commit) is required}"
|
||||
: "${PR_NUMBER:?PR_NUMBER (pull request number) is required}"
|
||||
: "${GITHUB_REPOSITORY:?GITHUB_REPOSITORY is required}"
|
||||
|
||||
# Workflow-level `name:` values that must be green before the full suite
|
||||
# may start. "Deploy Documentation" is path-filtered on PRs, so its run may
|
||||
# legitimately never exist; pre-commit always runs, so it must appear.
|
||||
WATCHED_NAMES='["pre-commit", "Deploy Documentation"]'
|
||||
WATCHED_REGEX='^(pre-commit|Deploy Documentation)$'
|
||||
POLL_SECS="${POLL_SECS:-20}"
|
||||
GRACE_SECS="${GRACE_SECS:-60}"
|
||||
MAX_WAIT_SECS="${MAX_WAIT_SECS:-1500}"
|
||||
|
||||
# Bound each API call so a hung connection hits the 3-strike fail-open path
|
||||
# instead of pinning the loop until the job timeout (which would fail closed
|
||||
# on exactly the GitHub-outage case this script is meant to survive).
|
||||
if command -v timeout >/dev/null 2>&1; then
|
||||
gh_api() { timeout 30 gh api "$@"; }
|
||||
else
|
||||
gh_api() { gh api "$@"; } # macOS dev boxes; CI always has coreutils timeout
|
||||
fi
|
||||
|
||||
# The workflow checked the label before starting the gate, but the wait can
|
||||
# last ~25 min: re-check once before any exit 0 and fail closed if 'ready'
|
||||
# was removed in the meantime. An API error here proceeds (the label was
|
||||
# present when the gate started; never brick CI on an outage).
|
||||
recheck_ready_label() {
|
||||
local pr_json
|
||||
if pr_json=$(gh_api "repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" 2>/dev/null); then
|
||||
if ! jq -e '[.labels[]?.name] | index("ready")' <<<"$pr_json" >/dev/null 2>&1; then
|
||||
echo "::error::PR #${PR_NUMBER} no longer has the 'ready' label —" \
|
||||
"NOT triggering the Buildkite full suite. Re-add the label to re-arm."
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
echo "::warning::Could not re-check the 'ready' label on PR #${PR_NUMBER}; proceeding (it was present when the gate started)."
|
||||
fi
|
||||
}
|
||||
|
||||
start=$(date +%s)
|
||||
api_fails=0
|
||||
missing=""
|
||||
|
||||
while true; do
|
||||
elapsed=$(( $(date +%s) - start ))
|
||||
|
||||
if runs_json=$(gh_api "repos/${GITHUB_REPOSITORY}/actions/runs?head_sha=${PR_SHA}&per_page=100" 2>/dev/null) \
|
||||
&& state=$(jq --arg re "$WATCHED_REGEX" '
|
||||
[.workflow_runs[]? | select(.name // "" | test($re))]
|
||||
| group_by(.name) | map(max_by(.id))
|
||||
| map({name, status, conclusion})' <<<"$runs_json" 2>/dev/null); then
|
||||
api_fails=0
|
||||
echo "t+${elapsed}s watched checks: $(jq -c . <<<"$state")"
|
||||
|
||||
failed=$(jq -r '[.[] | select(.status == "completed"
|
||||
and (.conclusion | IN("success", "skipped", "neutral", "cancelled") | not))]
|
||||
| map(.name) | join(", ")' <<<"$state")
|
||||
if [ -n "$failed" ]; then
|
||||
echo "::error::Cheap check(s) failed on ${PR_SHA}: ${failed}." \
|
||||
"NOT triggering the Buildkite full suite. Push a fix (the 'ready'" \
|
||||
"label re-arms on every push), or re-run the failed check and then" \
|
||||
"re-run this workflow."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 'cancelled' counts as pending: wait for a re-run to reach a real verdict
|
||||
# (bounded by MAX_WAIT, then the fail-open below).
|
||||
pending=$(jq '[.[] | select(.status != "completed" or .conclusion == "cancelled")] | length' <<<"$state")
|
||||
missing=$(jq -r --argjson watched "$WATCHED_NAMES" '($watched - map(.name)) | join(", ")' <<<"$state")
|
||||
if [ "$pending" -eq 0 ]; then
|
||||
if [ -z "$missing" ]; then
|
||||
recheck_ready_label
|
||||
echo "All watched cheap checks are green — full suite may proceed."
|
||||
exit 0
|
||||
fi
|
||||
case "$missing" in
|
||||
*pre-commit*)
|
||||
echo "pre-commit run not found for ${PR_SHA} yet; waiting (pre-commit is never path-filtered, so its absence is anomalous)."
|
||||
;;
|
||||
*)
|
||||
if [ "$elapsed" -ge "$GRACE_SECS" ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::Watched run(s) never appeared for ${PR_SHA}: ${missing} (path-filtered, likely not applicable). Proceeding on the checks that did run."
|
||||
exit 0
|
||||
fi
|
||||
echo "Waiting up to ${GRACE_SECS}s grace for path-filtered run(s) to appear: ${missing}."
|
||||
;;
|
||||
esac
|
||||
fi
|
||||
else
|
||||
api_fails=$(( api_fails + 1 ))
|
||||
echo "::warning::GitHub API error querying workflow runs for ${PR_SHA} (attempt ${api_fails}/3)."
|
||||
if [ "$api_fails" -ge 3 ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the full suite WITHOUT the cheap-check gate."
|
||||
exit 0
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$elapsed" -ge "$MAX_WAIT_SECS" ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the full suite anyway."
|
||||
exit 0
|
||||
fi
|
||||
sleep "$POLL_SECS"
|
||||
done
|
||||
Executable
+122
@@ -0,0 +1,122 @@
|
||||
#!/usr/bin/env bash
|
||||
# Self-test for gate_full_suite.sh using a mocked `gh`. No network, runs on
|
||||
# any dev box: bash .github/scripts/test_gate_full_suite.sh
|
||||
set -u
|
||||
here=$(cd "$(dirname "$0")" && pwd)
|
||||
tmp=$(mktemp -d)
|
||||
trap 'rm -rf "$tmp"' EXIT
|
||||
|
||||
# Mock gh. Asserts the exact endpoint (including head_sha) it is called
|
||||
# with — an endpoint typo in the gate script fails the test rather than
|
||||
# silently serving canned data. On the runs endpoint it serves
|
||||
# $MOCK_DIR/response_<call#>.json, sticking on the highest existing file,
|
||||
# and exits 1 if none exist (simulates a GitHub API outage). On the pulls
|
||||
# endpoint it serves $MOCK_DIR/pr.json, defaulting to a 'ready'-labeled PR.
|
||||
cat > "$tmp/gh" <<'EOF'
|
||||
#!/usr/bin/env bash
|
||||
if [ "${1:-}" != "api" ]; then
|
||||
echo "unexpected gh invocation: $*" >> "$MOCK_DIR/endpoint_error"
|
||||
exit 2
|
||||
fi
|
||||
case "${2:-}" in
|
||||
"repos/o/r/actions/runs?head_sha=deadbeef&per_page=100")
|
||||
n=$(( $(cat "$MOCK_DIR/count" 2>/dev/null || echo 0) + 1 ))
|
||||
echo "$n" > "$MOCK_DIR/count"
|
||||
while [ "$n" -gt 0 ]; do
|
||||
if [ -f "$MOCK_DIR/response_$n.json" ]; then
|
||||
cat "$MOCK_DIR/response_$n.json"
|
||||
exit 0
|
||||
fi
|
||||
n=$(( n - 1 ))
|
||||
done
|
||||
echo "api outage" >&2
|
||||
exit 1
|
||||
;;
|
||||
"repos/o/r/pulls/42")
|
||||
if [ -f "$MOCK_DIR/pr.json" ]; then
|
||||
cat "$MOCK_DIR/pr.json"
|
||||
else
|
||||
echo '{"labels": [{"name": "ready"}]}'
|
||||
fi
|
||||
;;
|
||||
*)
|
||||
echo "unexpected gh endpoint: $2" >> "$MOCK_DIR/endpoint_error"
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
EOF
|
||||
chmod +x "$tmp/gh"
|
||||
|
||||
PC_OK='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "success"}'
|
||||
PC_BAD='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "failure"}'
|
||||
PC_PENDING='{"name": "pre-commit", "id": 1, "status": "in_progress", "conclusion": null}'
|
||||
DOCS_OK='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "success"}'
|
||||
DOCS_BAD='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "failure"}'
|
||||
DOCS_CANCELLED='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "cancelled"}'
|
||||
OTHER='{"name": "Trigger Full Suite", "id": 3, "status": "in_progress", "conclusion": null}'
|
||||
NULL_NAME='{"name": null, "id": 4, "status": "completed", "conclusion": "failure"}'
|
||||
PC_OK_RERUN='{"name": "pre-commit", "id": 5, "status": "completed", "conclusion": "success"}'
|
||||
|
||||
fails=0
|
||||
want_log="" # optional: expect() also greps out.log for this regex, then resets
|
||||
pr_json="" # optional: served for the pulls (label re-check) endpoint, then resets
|
||||
raw_body="" # optional: serve responses verbatim instead of wrapping in workflow_runs
|
||||
expect() { # <name> <expected-exit> <response json>...
|
||||
local name=$1 want=$2 dir i=1
|
||||
shift 2
|
||||
dir=$(mktemp -d "$tmp/test_XXXXXX")
|
||||
for body in "$@"; do
|
||||
if [ -n "$raw_body" ]; then
|
||||
printf '%s' "$body" > "$dir/response_$i.json"
|
||||
else
|
||||
printf '{"workflow_runs": [%s]}' "$body" > "$dir/response_$i.json"
|
||||
fi
|
||||
i=$(( i + 1 ))
|
||||
done
|
||||
[ -n "$pr_json" ] && printf '%s' "$pr_json" > "$dir/pr.json"
|
||||
( export PATH="$tmp:$PATH" MOCK_DIR="$dir" PR_SHA=deadbeef PR_NUMBER=42 \
|
||||
GITHUB_REPOSITORY=o/r POLL_SECS=0 GRACE_SECS=1 MAX_WAIT_SECS=3
|
||||
bash "$here/gate_full_suite.sh" > "$dir/out.log" 2>&1 )
|
||||
local rc=$?
|
||||
if [ "$rc" -ne "$want" ]; then
|
||||
echo "FAIL: $name (exit $rc, want $want)"
|
||||
cat "$dir/out.log"
|
||||
fails=1
|
||||
elif [ -f "$dir/endpoint_error" ]; then
|
||||
echo "FAIL: $name (mock gh got an unexpected call)"
|
||||
cat "$dir/endpoint_error"
|
||||
fails=1
|
||||
elif [ -n "$want_log" ] && ! grep -Eq "$want_log" "$dir/out.log"; then
|
||||
echo "FAIL: $name (log does not match: $want_log)"
|
||||
cat "$dir/out.log"
|
||||
fails=1
|
||||
else
|
||||
echo "ok: $name"
|
||||
fi
|
||||
want_log="" pr_json="" raw_body=""
|
||||
}
|
||||
|
||||
expect "both green -> proceed" 0 "$PC_OK, $DOCS_OK, $OTHER, $NULL_NAME"
|
||||
expect "docs build failed -> blocked" 1 "$PC_OK, $DOCS_BAD"
|
||||
expect "pre-commit failed -> blocked" 1 "$PC_BAD"
|
||||
expect "pending then green -> proceed" 0 "$PC_PENDING" "$PC_OK, $DOCS_OK"
|
||||
want_log="never appeared.*Deploy Documentation"
|
||||
expect "docs run absent (path-filtered) -> proceed after grace" 0 "$PC_OK"
|
||||
expect "API outage -> fail open" 0
|
||||
want_log="FAILING OPEN"
|
||||
expect "pending past MAX_WAIT -> fail open" 0 "$PC_PENDING"
|
||||
want_log="FAILING OPEN"
|
||||
expect "unrelated runs only -> no grace, fail open at MAX_WAIT" 0 "$OTHER"
|
||||
expect "cancelled docs then green -> proceed" 0 \
|
||||
"$PC_OK, $DOCS_CANCELLED" "$PC_OK, $DOCS_OK"
|
||||
want_log="FAILING OPEN"
|
||||
expect "cancelled docs forever -> fail open at MAX_WAIT" 0 "$PC_OK, $DOCS_CANCELLED"
|
||||
want_log="FAILING OPEN"
|
||||
expect "pre-commit absent -> no grace, fail open at MAX_WAIT" 0 "$DOCS_OK"
|
||||
expect "duplicate run names -> latest wins" 0 "$PC_BAD, $PC_OK_RERUN, $DOCS_OK"
|
||||
raw_body=1
|
||||
expect "garbage response body -> fail open" 0 "this is not json"
|
||||
pr_json='{"labels": [{"name": "other"}]}'
|
||||
expect "ready label removed mid-gate -> blocked" 1 "$PC_OK, $DOCS_OK"
|
||||
|
||||
exit "$fails"
|
||||
@@ -1,7 +1,11 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
# pull_request_target instead of pull_request: the workflow definition and
|
||||
# the hook config are always taken from the BASE branch, so fork /
|
||||
# first-time-contributor PRs run immediately without a maintainer clicking
|
||||
# "Approve and run". The PR head is checked out as data only.
|
||||
pull_request_target:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
inputs:
|
||||
@@ -15,12 +19,25 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
if: github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
@@ -30,3 +47,6 @@ jobs:
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
# After pre-commit so a self-test failure cannot mask lint failures.
|
||||
- name: Full-suite gate self-test
|
||||
run: bash .github/scripts/test_gate_full_suite.sh
|
||||
|
||||
@@ -52,6 +52,7 @@ jobs:
|
||||
core.setOutput('pr_sha', pr.head.sha);
|
||||
core.setOutput('pr_branch', pr.head.ref);
|
||||
core.setOutput('pr_number', String(prNumber));
|
||||
core.setOutput('pr_title', pr.title);
|
||||
|
||||
- name: Trigger Full Suite
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
@@ -60,6 +61,7 @@ jobs:
|
||||
PR_SHA: ${{ steps.label.outputs.pr_sha }}
|
||||
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
|
||||
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
|
||||
PR_TITLE: ${{ steps.label.outputs.pr_title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -71,6 +73,7 @@ jobs:
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
@@ -80,11 +83,12 @@ jobs:
|
||||
pull_request_id: $pr_id,
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
}
|
||||
}')"
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring),
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
|
||||
parse-command:
|
||||
if: >-
|
||||
@@ -125,7 +129,7 @@ jobs:
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
@@ -136,6 +140,7 @@ jobs:
|
||||
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
@@ -240,6 +245,7 @@ jobs:
|
||||
TEST_SCOPE: ${{ needs.parse-command.outputs.test_scope }}
|
||||
FULL_SUITE: ${{ needs.parse-command.outputs.full_suite }}
|
||||
TEST_TYPE: ${{ needs.parse-command.outputs.test_type }}
|
||||
PR_TITLE: ${{ github.event.issue.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -256,6 +262,7 @@ jobs:
|
||||
--arg full_suite "$FULL_SUITE" \
|
||||
--arg test_type "$TEST_TYPE" \
|
||||
--arg pr_number "$PR_NUMBER" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
'{
|
||||
commit: $commit,
|
||||
branch: $branch,
|
||||
@@ -265,8 +272,9 @@ jobs:
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: $test_scope,
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number
|
||||
}
|
||||
}')"
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number,
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
|
||||
@@ -7,6 +7,7 @@ on:
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
actions: read
|
||||
|
||||
concurrency:
|
||||
group: full-suite-${{ github.event.pull_request.number }}
|
||||
@@ -18,6 +19,8 @@ jobs:
|
||||
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|
||||
|| github.event.action == 'synchronize'
|
||||
runs-on: ubuntu-latest
|
||||
# Gate below may wait for cheap checks (up to MAX_WAIT_SECS = 25 min).
|
||||
timeout-minutes: 35
|
||||
steps:
|
||||
- name: Check ready label
|
||||
id: check
|
||||
@@ -49,6 +52,20 @@ jobs:
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
|
||||
done
|
||||
|
||||
# Checks out the BASE branch (default for pull_request_target), so PR
|
||||
# authors cannot tamper with the gate script.
|
||||
- name: Checkout gate script
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
|
||||
|
||||
- name: Wait for pre-commit and docs build
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
run: bash .github/scripts/gate_full_suite.sh
|
||||
|
||||
- name: Trigger Buildkite Full Suite
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
@@ -56,6 +73,7 @@ jobs:
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_TITLE: ${{ github.event.pull_request.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -67,6 +85,7 @@ jobs:
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER}" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
@@ -78,6 +97,7 @@ jobs:
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
PR_NUMBER: ($pr_id | tostring),
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
|
||||
@@ -13,12 +13,33 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
|
||||
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
|
||||
# push trigger is a sufficient change detector on its own -- no separate
|
||||
# detect-changes/paths-filter job is needed now that there is a single
|
||||
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
|
||||
# rocm Dockerfile stay manual-dispatch only.
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'docker/Dockerfile'
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
# One static group, no cancellation: every run of this workflow writes the same
|
||||
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
|
||||
# concurrent push/dispatch runs would race on those tags, and cancelling a run
|
||||
# mid-publish can strand the cu126/cu130 tag families at different commits. An
|
||||
# in-flight superseded build wastes its runner time, but its tags are then
|
||||
# overwritten by the newer queued run. GitHub keeps a single pending run per
|
||||
# group: the newest queued run replaces any older queued one.
|
||||
concurrency:
|
||||
group: infra-build-image
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
|
||||
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
|
||||
@@ -28,7 +49,11 @@ jobs:
|
||||
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
|
||||
# 2.8.3 comes from the architecture-specific prebuilt releases.
|
||||
build-cuda-images:
|
||||
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
|
||||
# on a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# repository guard keeps fork syncs from auto-building; manual dispatch
|
||||
# still works in forks.
|
||||
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -75,10 +100,10 @@ jobs:
|
||||
secrets: inherit
|
||||
|
||||
publish-cuda-manifests:
|
||||
# !cancelled(): a failed sibling build leg must not skip the manifests for a
|
||||
# CUDA lane whose own digests all exist; the digest-count check below fails
|
||||
# the incomplete lane loudly instead.
|
||||
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# !cancelled(): publish lanes whose digests exist even if a sibling build
|
||||
# leg failed (the digest-count check fails incomplete lanes); it also
|
||||
# bypasses skipped-needs propagation, hence the explicit skipped check.
|
||||
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
|
||||
needs: build-cuda-images
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
|
||||
+3
-2
@@ -72,8 +72,7 @@ docs/distillation/examples/
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
# Reference videos (negations must come after the catch-all on line below)
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
@@ -127,6 +126,8 @@ apps/dreamverse/web/.env.production.local
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.png
|
||||
|
||||
# Editor logs and local Python version pins (accidentally committed)
|
||||
*.nvimlog
|
||||
|
||||
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
|
||||
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
|
||||
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
|
||||
|
||||
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
|
||||
# the dreamverse extra and is validated with it, so enable it here.
|
||||
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
|
||||
STREAM_MODE=av_fmp4 \
|
||||
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
|
||||
FASTVIDEO_FA4=1 \
|
||||
HF_HOME=/root/.cache/huggingface
|
||||
|
||||
RUN mkdir -p /var/lib/dreamverse
|
||||
|
||||
@@ -12,13 +12,17 @@ Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata:
|
||||
Records can include source metadata and rolling-baseline policy context:
|
||||
|
||||
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
|
||||
- `baseline_eligible`: only successful scheduled-main records should be true
|
||||
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
|
||||
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
|
||||
floors used for recomputed status context
|
||||
|
||||
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
|
||||
crossings; `regressed` remains the gated CI-failure signal.
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
- PR number, branch, and Buildkite URL when present
|
||||
|
||||
The latest status table uses the stored JSON `success` value. Recomputed
|
||||
baseline context is shown separately and does not override stored status.
|
||||
baseline context applies each metric's percent and absolute regression floors
|
||||
and does not override stored status.
|
||||
|
||||
## API
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
||||
import { fetchSummary, fetchTrends, refreshData, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
import { fetchSummary, fetchTrends, refreshData } from "./api";
|
||||
import type { CohortValue, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
|
||||
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
|
||||
const RUN_SOURCES: Array<{ value: "" | RunSource; label: string }> = [
|
||||
@@ -109,6 +110,61 @@ function metricLabel(metricKey: string) {
|
||||
return METRIC_DEFINITIONS[metricKey]?.label ?? metricKey;
|
||||
}
|
||||
|
||||
type CohortFields = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
workload_id: CohortValue;
|
||||
variant_id: CohortValue;
|
||||
benchmark_version: CohortValue;
|
||||
recipe_fingerprint: CohortValue;
|
||||
hardware_profile_id: CohortValue;
|
||||
software_profile_id: CohortValue;
|
||||
};
|
||||
|
||||
function cohortValue(value: CohortValue) {
|
||||
if (value === null || value === undefined || value === "") {
|
||||
return "legacy";
|
||||
}
|
||||
return String(value);
|
||||
}
|
||||
|
||||
function shortCohortValue(value: CohortValue) {
|
||||
const text = cohortValue(value);
|
||||
if (text === "legacy" || text.length <= 14) {
|
||||
return text;
|
||||
}
|
||||
return text.slice(0, 12);
|
||||
}
|
||||
|
||||
function cohortKey(cohort: CohortFields) {
|
||||
return [
|
||||
cohort.model_id,
|
||||
cohort.gpu_type,
|
||||
cohortValue(cohort.workload_id),
|
||||
cohortValue(cohort.variant_id),
|
||||
cohortValue(cohort.benchmark_version),
|
||||
cohortValue(cohort.recipe_fingerprint),
|
||||
cohortValue(cohort.hardware_profile_id),
|
||||
cohortValue(cohort.software_profile_id)
|
||||
].join("|");
|
||||
}
|
||||
|
||||
function cohortTitle(cohort: CohortFields) {
|
||||
const workload = cohortValue(cohort.workload_id);
|
||||
const variant = cohortValue(cohort.variant_id);
|
||||
const version = cohortValue(cohort.benchmark_version);
|
||||
const versionLabel = version === "legacy" ? version : `v${version}`;
|
||||
return `${workload} / ${variant} / ${versionLabel}`;
|
||||
}
|
||||
|
||||
function cohortDetail(cohort: CohortFields) {
|
||||
return [
|
||||
`recipe ${shortCohortValue(cohort.recipe_fingerprint)}`,
|
||||
shortCohortValue(cohort.hardware_profile_id),
|
||||
shortCohortValue(cohort.software_profile_id)
|
||||
].join(" | ");
|
||||
}
|
||||
|
||||
function formatMetricValue(metricKey: string, value: number | null | undefined, tooltip = false) {
|
||||
const definition = METRIC_DEFINITIONS[metricKey];
|
||||
if (!definition) {
|
||||
@@ -171,7 +227,9 @@ function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string
|
||||
top: `${(activePoint.y / height) * 100}%`
|
||||
}
|
||||
: undefined;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}`;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}, ${cohortTitle(
|
||||
group
|
||||
)}`;
|
||||
|
||||
return (
|
||||
<div className="chart-shell">
|
||||
@@ -419,7 +477,7 @@ export default function App() {
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Latest Status</h2>
|
||||
<span>{latestRows.length} model/GPU groups</span>
|
||||
<span>{latestRows.length} comparison cohorts</span>
|
||||
</div>
|
||||
{latestRows.length === 0 ? (
|
||||
<div className="empty">No records match the selected filters.</div>
|
||||
@@ -432,6 +490,7 @@ export default function App() {
|
||||
<th>Recomputed</th>
|
||||
<th>Model</th>
|
||||
<th>GPU</th>
|
||||
<th>Cohort</th>
|
||||
<th>Commit</th>
|
||||
<th>Source</th>
|
||||
<th>Baseline</th>
|
||||
@@ -440,11 +499,13 @@ export default function App() {
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
<th>Exceeded</th>
|
||||
<th>Failing</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{latestRows.map((row) => (
|
||||
<tr key={`${row.model_id}-${row.gpu_type}`}>
|
||||
<tr key={cohortKey(row)}>
|
||||
<td>
|
||||
<span className={`badge ${row.status}`}>{row.status}</span>
|
||||
</td>
|
||||
@@ -455,6 +516,12 @@ export default function App() {
|
||||
</td>
|
||||
<td>{row.model_id}</td>
|
||||
<td>{row.gpu_type}</td>
|
||||
<td>
|
||||
<div className="cohort-cell">
|
||||
<strong>{cohortTitle(row)}</strong>
|
||||
<span>{cohortDetail(row)}</span>
|
||||
</div>
|
||||
</td>
|
||||
<td>{shortSha(row.commit_sha)}</td>
|
||||
<td>
|
||||
<span className={`source-badge source-${row.run_source}`}>{runSourceLabel(row.run_source)}</span>
|
||||
@@ -465,6 +532,12 @@ export default function App() {
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
<td>
|
||||
{row.threshold_exceeded_metrics.length
|
||||
? row.threshold_exceeded_metrics.join(", ")
|
||||
: "none"}
|
||||
</td>
|
||||
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
@@ -487,11 +560,13 @@ export default function App() {
|
||||
) : (
|
||||
trends.map((group) =>
|
||||
METRIC_KEYS.map((metricKey) => (
|
||||
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
|
||||
<article className="trend-card" key={`${cohortKey(group)}-${metricKey}`}>
|
||||
<div>
|
||||
<h3>{metricLabel(metricKey)}</h3>
|
||||
<p>
|
||||
{group.model_id} | {group.gpu_type}
|
||||
<span>{cohortTitle(group)}</span>
|
||||
<span>{cohortDetail(group)}</span>
|
||||
</p>
|
||||
</div>
|
||||
<TrendChart group={group} metricKey={metricKey} />
|
||||
|
||||
@@ -2,11 +2,28 @@ export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
absolute_delta: number | null;
|
||||
threshold_percent: number;
|
||||
threshold_absolute: number;
|
||||
gated: boolean;
|
||||
threshold_exceeded: boolean;
|
||||
regressed: boolean;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
};
|
||||
|
||||
export type CohortValue = string | number | null;
|
||||
|
||||
export type ComparisonCohort = {
|
||||
workload_id: CohortValue;
|
||||
variant_id: CohortValue;
|
||||
benchmark_version: CohortValue;
|
||||
recipe_fingerprint: CohortValue;
|
||||
hardware_profile_id: CohortValue;
|
||||
software_profile_id: CohortValue;
|
||||
};
|
||||
|
||||
export type SummaryRow = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
@@ -15,7 +32,8 @@ export type SummaryRow = {
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
regression_threshold_pct: number;
|
||||
threshold_exceeded_metrics: string[];
|
||||
failing_metrics: string[];
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
@@ -27,7 +45,7 @@ export type SummaryRow = {
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, MetricValue>;
|
||||
};
|
||||
} & ComparisonCohort;
|
||||
|
||||
export type RunSource = "pr" | "local" | "scheduled_main" | "unknown";
|
||||
|
||||
@@ -61,13 +79,13 @@ export type TrendPoint = {
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, number | null>;
|
||||
};
|
||||
} & ComparisonCohort;
|
||||
|
||||
export type TrendGroup = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
points: TrendPoint[];
|
||||
};
|
||||
} & ComparisonCohort;
|
||||
|
||||
export type TrendsResponse = {
|
||||
groups: TrendGroup[];
|
||||
|
||||
@@ -149,6 +149,11 @@ h3 {
|
||||
font-size: 0.82rem;
|
||||
}
|
||||
|
||||
.trend-card p {
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.stat strong {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
@@ -186,7 +191,7 @@ h3 {
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
min-width: 1120px;
|
||||
min-width: 1260px;
|
||||
border-collapse: collapse;
|
||||
}
|
||||
|
||||
@@ -209,6 +214,25 @@ td {
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.cohort-cell {
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.cohort-cell strong,
|
||||
.trend-card p span {
|
||||
color: #1b2836;
|
||||
font-size: 0.78rem;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.cohort-cell span,
|
||||
.trend-card p span + span {
|
||||
color: #607080;
|
||||
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", monospace;
|
||||
font-size: 0.72rem;
|
||||
}
|
||||
|
||||
.badge {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
{
|
||||
"alpha_yaw": 0.08734091699186919,
|
||||
"alpha_pitch": 0.08169667696275307,
|
||||
"alpha_turn": 5.724587470723463e-17,
|
||||
"beta_fwd": 0.02842768078408099,
|
||||
"beta_strafe": 0.022531015077067108,
|
||||
"focal_length": 457.0,
|
||||
"frame_shape": [
|
||||
352,
|
||||
640
|
||||
],
|
||||
"calibrated_from": [
|
||||
"1_wasd_only",
|
||||
"camera",
|
||||
"camera4hold_alpha1",
|
||||
"fully_random",
|
||||
"wasdonly_alpha1",
|
||||
"wasd4holdrandview_simple_1key1mouse1"
|
||||
],
|
||||
"residual_rms": 15.890399609478676,
|
||||
"n_equations": 4125232
|
||||
}
|
||||
+6
-4
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
|
||||
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
|
||||
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
|
||||
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
|
||||
ARG FA4_CUTE_REF=940cd9680f3315f2f06b43ab5bea2c2cf2d96806
|
||||
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
|
||||
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
|
||||
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
|
||||
@@ -161,7 +161,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
uv pip install flash-attn==${FLASH_ATTN_VERSION} --no-build-isolation; \
|
||||
fi
|
||||
|
||||
# Overlay the cutlass-4.5-safe upstream FA4 cute (FA4_CUTE_REF) over the
|
||||
# Overlay the CuTe-DSL-4.6-compatible upstream FA4 cute (FA4_CUTE_REF) over the
|
||||
# wheel/source one so the image runs FA4, not the FA2 fallback. This pulls the FA4
|
||||
# runtime stack (cutlass-dsl, quack-kernels, apache-tvm-ffi, torch-c-dlpack-ext) --
|
||||
# the same deps the [dreamverse] extra already installs in CI; the installed torch
|
||||
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
|
||||
# Then verify both survive so a broken overlay fails the build instead of shipping
|
||||
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
|
||||
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
|
||||
# arm64 and the image runs FA3/FA2 as usual.
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
|
||||
else \
|
||||
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
|
||||
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
|
||||
|
||||
@@ -99,10 +99,18 @@ status.
|
||||
Full Suite is also path-filtered. It validates broader behavior before Mergify
|
||||
can merge a PR.
|
||||
|
||||
A `ready`-labeled PR does not hit Buildkite immediately:
|
||||
`ci-trigger-full-suite.yml` first runs `.github/scripts/gate_full_suite.sh`,
|
||||
which waits for the cheap Tier-1 checks (pre-commit, docs build) on the PR
|
||||
head. A red cheap check blocks the suite (fail closed; the next push re-arms
|
||||
it), while a GitHub outage or a >25 min wait lets it run anyway (fail open).
|
||||
`/test full` bypasses the gate.
|
||||
|
||||
| Buildkite label | `TEST_TYPE` | Main watched paths |
|
||||
|---|---|---|
|
||||
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
|
||||
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
|
||||
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
|
||||
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
|
||||
@@ -144,6 +152,7 @@ Valid direct test names:
|
||||
| `/test training` | `training` |
|
||||
| `/test lora-inference` | `inference_lora` |
|
||||
| `/test lora-training` | `training_lora` |
|
||||
| `/test lora-extraction` | `lora_extraction` |
|
||||
| `/test distillation` | `distillation_dmd` |
|
||||
| `/test self-forcing` | `self_forcing` |
|
||||
| `/test vsa` | `training_vsa` |
|
||||
|
||||
@@ -72,14 +72,18 @@ fastvideo/tests/performance/
|
||||
│ writes Markdown summary + (optionally) uploads new records
|
||||
├── dashboard.py
|
||||
│ └── builds time-series Plotly HTML from HF history
|
||||
└── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
|
||||
fastvideo/performance/
|
||||
├── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
└── metric_policy.py # shared rolling-baseline threshold policy
|
||||
```
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
normalized JSON per `(model_id, gpu_type, run)` tuple. The rolling baseline is
|
||||
the median of the last 5 successful, baseline-eligible records for that
|
||||
model+GPU. PR and local records are visible in the dashboard but are not
|
||||
baseline eligible.
|
||||
normalized JSON per run. For v2 records, the rolling baseline is the median of
|
||||
the last 5 successful, baseline-eligible records in the same comparison cohort:
|
||||
`model_id`, `gpu_type`, `workload_id`, `variant_id`, `benchmark_version`,
|
||||
`recipe_fingerprint`, `hardware_profile_id`, and `software_profile_id`. PR and
|
||||
local records are visible in the dashboard but are not baseline eligible.
|
||||
|
||||
## Planned Coverage
|
||||
|
||||
@@ -92,25 +96,28 @@ and recipe changes instead of treating all records for a model as equivalent.
|
||||
|
||||
## Metrics
|
||||
|
||||
Each benchmark records six metrics:
|
||||
Each benchmark records six metrics. The rolling-baseline comparator also has a
|
||||
per-metric policy with direction, percent threshold, absolute threshold, and a
|
||||
`gated` flag.
|
||||
|
||||
| Metric | Raw key | Normalized key | Direction |
|
||||
|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
|
||||
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|
||||
|---|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
|
||||
|
||||
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
|
||||
while it runs so pipeline stage execution times are available in
|
||||
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
|
||||
`prompt_encoding_stage` so duplicate stage classes do not collide. For
|
||||
`PipelineStage` entries, the extractor maps the `stage_class` field:
|
||||
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
|
||||
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
|
||||
`vae_decode_time_s`, with a fallback for older logs that used the class name as
|
||||
`PipelineStage` entries, shared component stage bases emit a stable
|
||||
`component_metric`: text encoding stages map to `text_encoder_time_s`,
|
||||
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
|
||||
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
|
||||
older logs that do not include `component_metric` or that used the class name as
|
||||
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
|
||||
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
|
||||
does not report one of the mapped stages, that component metric is stored as
|
||||
@@ -152,13 +159,29 @@ unrealistic memory growth, and optionally large component-specific slowdowns
|
||||
even when the rolling baseline is empty. They are hand-set with generous
|
||||
headroom and almost never need touching.
|
||||
|
||||
### Rolling baseline (per `(model_id, gpu_type)`)
|
||||
### Rolling baseline (per comparison cohort)
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
|
||||
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
|
||||
for each available metric, and fails if the current run regresses by more than
|
||||
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
|
||||
higher values are regressions. For throughput, lower values are regressions.
|
||||
for the same comparison cohort from the HF dataset, computes the median for
|
||||
each available metric, and evaluates the current run with the metric's
|
||||
rolling regression policy. For v2 records, that cohort is `model_id`,
|
||||
`gpu_type`, `workload_id`, `variant_id`, `benchmark_version`,
|
||||
`recipe_fingerprint`, `hardware_profile_id`, and `software_profile_id`. For
|
||||
latency, memory, and component times, higher values are regressions. For
|
||||
throughput, lower values are regressions.
|
||||
|
||||
A metric exceeds its rolling threshold when both of these are true:
|
||||
|
||||
```text
|
||||
percent_delta > threshold_percent
|
||||
absolute_delta > threshold_absolute
|
||||
```
|
||||
|
||||
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
|
||||
for metrics that should remain visible in reports and the dashboard without
|
||||
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
|
||||
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
|
||||
metrics are skipped.
|
||||
|
||||
This is the **drift detector** — it catches sub-threshold regressions that
|
||||
slowly add up. Only scheduled-main successful records are baseline eligible.
|
||||
@@ -172,6 +195,51 @@ agent skill to advance the rolling median.
|
||||
|
||||
## Schemas
|
||||
|
||||
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
|
||||
|
||||
Benchmark configs without `config_schema_version` are treated as legacy v1
|
||||
configs and remain loadable. New or migrated configs should use
|
||||
`config_schema_version: 2` and include explicit comparable identity fields:
|
||||
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2
|
||||
}
|
||||
```
|
||||
|
||||
`benchmark_id` is still required in this phase because raw artifact names,
|
||||
generated-video directories, normalized record paths, and the current rolling
|
||||
baseline comparator still depend on it. The v2 identity fields are config
|
||||
metadata that make the measured workload explicit:
|
||||
|
||||
| Field | Purpose |
|
||||
|---|---|
|
||||
| `workload_id` | Stable benchmark family, such as `wan-t2v`. |
|
||||
| `variant_id` | Intentional recipe family, including model size and parallelism config, such as `1.3b-sp2`. |
|
||||
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
|
||||
|
||||
If a config declares `config_schema_version: 2`, loading fails clearly when any
|
||||
required v2 identity field is missing. If v2 identity or metadata fields are
|
||||
added without `config_schema_version: 2`, loading also fails so partial
|
||||
migrations do not silently run as v1 configs. Optional v2 metadata fields
|
||||
reserved for follow-up work, such as `metric_threshold_policy` and
|
||||
`quality_metadata`, must be JSON objects when present. (`recipe` is emitted
|
||||
by the harness and is not config-declarable.)
|
||||
|
||||
Recipe fingerprinting, hardware/software profile IDs, exact-identity
|
||||
comparison, and dashboard cohort grouping land with this change: v2 records
|
||||
compare only within their identity cohort, and a record that opens a NEW
|
||||
cohort is marked `baseline_status: "initialized_new_cohort"` (regression
|
||||
gating starts once that cohort accumulates history). Legacy v1 configs still
|
||||
run and are normalized for reporting, but their records skip rolling-baseline
|
||||
comparison entirely (`baseline_status: "skipped_missing_identity"`, never
|
||||
baseline eligible); only static thresholds gate them. Metric-specific
|
||||
threshold policies and promoted baselines remain separate follow-ups.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
@@ -179,6 +247,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -196,12 +268,62 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"max_dit_time_s": 10.0,
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 1.0,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"commit": "<full sha>",
|
||||
"run_source": "pr",
|
||||
"branch": "feature/perf-change",
|
||||
"pr_number": "1234",
|
||||
"test_scope": "direct",
|
||||
"build_url": "https://buildkite.example/build",
|
||||
"build_id": "<buildkite-build-id>",
|
||||
"job_id": "<buildkite-job-id>",
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"quality_metadata": { "quality_status": "canonical" },
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208
|
||||
"vae_decode_time_s": 3.208,
|
||||
"recipe": {
|
||||
"recipe_schema_version": 1,
|
||||
"benchmark": {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2
|
||||
},
|
||||
"model": { "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" },
|
||||
"init_kwargs": { "num_gpus": 2, "sp_size": 2, "tp_size": 1 },
|
||||
"generation_kwargs": { "height": 480, "width": 832, "num_frames": 45 },
|
||||
"inputs": { "prompt_count": 1, "prompt_sha256": ["<measured-prompt-sha256>"] },
|
||||
"attention": { "requested_backend": "FLASH_ATTN", "resolved_backend": "FLASH_ATTN" }
|
||||
},
|
||||
"recipe_fingerprint": "<sha256>",
|
||||
"hardware_profile": {
|
||||
"device_type": "cuda",
|
||||
"gpu_count": 2,
|
||||
"gpus": [{ "name": "NVIDIA L40S", "memory_gb": 48, "compute_capability": "8.9" }],
|
||||
"interconnect": "none_or_partial"
|
||||
},
|
||||
"hardware_profile_id": "hw-<sha256-prefix>",
|
||||
"software_profile": {
|
||||
"python": "3.12",
|
||||
"pytorch": "2.12",
|
||||
"cuda": "13.0",
|
||||
"packages": {
|
||||
"fastvideo_kernel": "0.3.2",
|
||||
"flashinfer": "0.2.11",
|
||||
"nvidia_cutlass_dsl": "4.5.0",
|
||||
"triton": "3.4.1"
|
||||
}
|
||||
},
|
||||
"software_profile_id": "sw-<sha256-prefix>",
|
||||
"environment_metadata": { "env": { "IMAGE_VERSION": "py3.12-cuda13.0.0" } },
|
||||
"environment_fingerprint": "env-<sha256-prefix>"
|
||||
}
|
||||
```
|
||||
|
||||
@@ -213,6 +335,10 @@ result, used as the rolling-baseline source of truth.
|
||||
```jsonc
|
||||
{
|
||||
"model_id": "wan-t2v-1.3b-2gpu",
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2,
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"commit_sha": "<full sha>",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
@@ -222,34 +348,69 @@ result, used as the rolling-baseline source of truth.
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.08,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"recipe_fingerprint": "<sha256>",
|
||||
"hardware_profile_id": "hw-<sha256-prefix>",
|
||||
"software_profile_id": "sw-<sha256-prefix>",
|
||||
"environment_fingerprint": "env-<sha256-prefix>",
|
||||
"run_source": "pr",
|
||||
"branch": "feature/perf-change",
|
||||
"pr_number": "1234",
|
||||
"test_scope": "direct",
|
||||
"build_url": "https://buildkite.example/build",
|
||||
"build_id": "<buildkite-build-id>",
|
||||
"job_id": "<buildkite-job-id>",
|
||||
"quality_metadata": { "quality_status": "canonical" },
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
### Compatibility with legacy records
|
||||
|
||||
Older records in the HF dataset may not have component timing fields. The
|
||||
comparator ignores missing or `null` metrics when computing a median, and the
|
||||
dashboard lists skipped plots for metric series that have no non-null values.
|
||||
Records missing both `run_source` and `baseline_eligible` are treated as legacy
|
||||
successful main/full-suite uploads and remain eligible for rolling baselines.
|
||||
Older records in the HF dataset may not have `result_schema_version`,
|
||||
component timing fields, or v2 identity/profile fields. Records without
|
||||
`result_schema_version` are treated as v1. The comparator ignores missing or
|
||||
`null` metrics when computing a median, and the dashboard lists skipped plots
|
||||
for metric series that have no non-null values. Records missing both
|
||||
`run_source` and `baseline_eligible` are treated as legacy successful
|
||||
main/full-suite uploads and remain eligible for rolling baselines.
|
||||
Current `perf_*.json` artifacts that lack the v2 comparison identity are
|
||||
normalized for reporting but skip rolling-baseline comparison and are not marked
|
||||
baseline eligible.
|
||||
|
||||
New records compare only against the same `model_id`, `gpu_type`,
|
||||
`workload_id`, `variant_id`, `benchmark_version`, `recipe_fingerprint`,
|
||||
`hardware_profile_id`, and `software_profile_id` cohort.
|
||||
`environment_metadata` and `environment_fingerprint` are audit data and are not
|
||||
part of the comparison key.
|
||||
The recipe prompt digests describe the prompts actually measured by the
|
||||
benchmark run; extra configured prompts are ignored unless the benchmark runner
|
||||
executes them.
|
||||
Software profile package cohorts keep exact versions for relevant
|
||||
attention/kernel packages, including FastVideo kernels, FlashAttention,
|
||||
FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
|
||||
|
||||
## Environment variable reference
|
||||
|
||||
| Variable | Default | Used by | Purpose |
|
||||
|---|---|---|---|
|
||||
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
|
||||
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
|
||||
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py`, `test_inference_performance.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
|
||||
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
|
||||
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
|
||||
|
||||
## CI integration
|
||||
@@ -260,17 +421,22 @@ point is `fastvideo/tests/modal/pr_test.py:run_performance_tests` and the
|
||||
Buildkite artifact upload is in
|
||||
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
|
||||
|
||||
Each performance build runs pytest first. If that fixed-threshold phase fails,
|
||||
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
|
||||
artifacts are not emitted. The dashboard still runs best-effort for
|
||||
observability. When pytest passes, the rolling-baseline phase emits:
|
||||
Each performance build runs pytest first. PR and direct runs only continue to
|
||||
`compare_baseline.py` when that fixed-threshold phase passes; if pytest fails,
|
||||
Markdown summaries and normalized JSON artifacts are not emitted. Scheduled
|
||||
main runs set `PERF_UPLOAD_POLICY=always`, so they still run
|
||||
`compare_baseline.py` (with `PERF_PYTEST_RC` set) after a fixed-threshold
|
||||
failure. Those failed scheduled main runs emit summaries and normalized
|
||||
records, upload records with `success=false`, and are excluded from future
|
||||
rolling baselines. The dashboard still runs best-effort for observability.
|
||||
When the rolling-baseline phase runs, it emits:
|
||||
|
||||
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
|
||||
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
|
||||
per-benchmark row with current vs. baseline values for latency, throughput,
|
||||
memory, text encoder time, DiT time, and VAE decode time.
|
||||
* **Plotly dashboard** — `dashboard_<sha>_<ts>.html` showing time-series for
|
||||
each metric grouped by `(model_id, gpu_type)`.
|
||||
each metric grouped by comparison cohort.
|
||||
* **Normalized records** — `normalized_perf_*.json`, one per benchmark.
|
||||
Useful as input to the
|
||||
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
|
||||
@@ -279,11 +445,16 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
## Adding a new benchmark
|
||||
|
||||
1. Drop a new JSON config into
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
|
||||
use v2 identity fields:
|
||||
|
||||
```json
|
||||
{
|
||||
"benchmark_id": "<unique-id>",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "<stable-workload-id>",
|
||||
"variant_id": "<variant, e.g. 1.3b-sp2>",
|
||||
"benchmark_version": 1,
|
||||
"model": { "model_path": "...", "model_short_name": "..." },
|
||||
"init_kwargs": { "num_gpus": 1, ... },
|
||||
"generation_kwargs": { "num_frames": 45, ... },
|
||||
@@ -299,9 +470,19 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
|
||||
}
|
||||
}
|
||||
```
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
Legacy v1 configs without `config_schema_version` still load, but should not
|
||||
gain v2 identity or metadata fields until they are migrated to
|
||||
`config_schema_version: 2`. For v2 configs, `workload_id`, `variant_id`,
|
||||
and `benchmark_version` are part of the comparison key; benchmark runs
|
||||
fail if any of these identity fields are missing.
|
||||
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
@@ -320,10 +501,17 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
a useful fixed gate. The rolling baseline will still track component times
|
||||
when static component thresholds are omitted.
|
||||
|
||||
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
|
||||
include only benchmark-specific deviations. Tune these independently from
|
||||
the fixed thresholds when a metric is noisy or should be informational. The
|
||||
fixed `thresholds` block is an absolute pytest ceiling. The
|
||||
`regression_thresholds` block controls rolling-baseline comparisons against
|
||||
recent scheduled-main records.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**"No baseline for ... Initializing"** — first run for this `(model_id,
|
||||
gpu_type)`. Run will pass and (if persisting) seed the first record.
|
||||
**"No baseline for ... Initializing"** — first run for this comparison cohort.
|
||||
Run will pass and (if persisting) seed the first record.
|
||||
|
||||
**Persistent failure right after a torch / kernel / image upgrade** —
|
||||
genuine regression *or* baseline drift. Compare the failing normalized record
|
||||
@@ -336,5 +524,5 @@ pipelines that did not report a mapped component stage.
|
||||
|
||||
**Component timing is `null`** — the generated result did not include a mapped
|
||||
stage in `logging_info.stages`. Check that the pipeline emits stage logging
|
||||
and that the stage name is listed in `STAGE_METRIC_MAP` in
|
||||
`test_inference_performance.py`.
|
||||
and that the stage emits `component_metric` or is covered by the legacy
|
||||
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
|
||||
|
||||
@@ -180,6 +180,30 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--device-folder L40S_reference_videos
|
||||
```
|
||||
|
||||
### SSIM Bootstrap Mode
|
||||
|
||||
Normal SSIM runs are strict: if a reference video or latent is missing, the
|
||||
test fails. For new-model PRs, CI can run SSIM in bootstrap mode so missing
|
||||
references are uploaded as draft artifacts for review instead of immediately
|
||||
blocking on a missing canonical reference.
|
||||
|
||||
Buildkite enables SSIM bootstrap mode when either condition is true:
|
||||
|
||||
- the PR title or Buildkite message contains `[new-model]`;
|
||||
- `FASTVIDEO_SSIM_BOOTSTRAP_MODE=1` is set for the Buildkite job.
|
||||
|
||||
Bootstrap mode passes `--ssim-bootstrap-mode` to pytest. When a generated
|
||||
artifact is available, the test uploads it under the `drafts/...` namespace in
|
||||
the SSIM reference repo and marks that case as expected-failed. After reviewing
|
||||
the draft, promote it into the canonical reference layout:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py promote-draft \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id <model_id>
|
||||
```
|
||||
|
||||
## CI Integration
|
||||
|
||||
FastVideo CI tests are orchestrated by Buildkite and run on Modal GPU
|
||||
|
||||
@@ -191,6 +191,9 @@ surfaces:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
color_correction_strength:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.dreamx_world.DreamXWorld5BARPipelineConfig
|
||||
default_camera_rotation:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
@@ -455,6 +458,8 @@ surfaces:
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
use_embedded_guidance: request.sampling.use_embedded_guidance
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# 🌊 AnyFlow Any-Step Video Distillation
|
||||
|
||||
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
|
||||
|
||||
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
|
||||
|
||||
```
|
||||
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
|
||||
```
|
||||
|
||||
for any ``t > r``.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
|
||||
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
|
||||
|
||||
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
|
||||
|
||||
## ⚙️ Inference
|
||||
|
||||
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
|
||||
|
||||
## 🧠 Algorithm
|
||||
|
||||
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
|
||||
|
||||
### Stage 1 — Pretrain (flow-map central-difference)
|
||||
|
||||
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
|
||||
|
||||
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
|
||||
|
||||
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
|
||||
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
|
||||
- the remainder is free.
|
||||
|
||||
The student forward at ``(t, r)`` is trained against the central-difference target
|
||||
|
||||
```
|
||||
target = (eps - x_0) - (t - r) · dF/dt
|
||||
```
|
||||
|
||||
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
|
||||
|
||||
### Stage 2 — On-policy DMD
|
||||
|
||||
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
|
||||
|
||||
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
|
||||
|
||||
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
### Stage 1 — pretrain
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
|
||||
```
|
||||
|
||||
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
|
||||
|
||||
- Global batch size: 32 (8 GPUs × 4 per-GPU)
|
||||
- Learning rate: 5e-5
|
||||
- Flow shift: 5.0
|
||||
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
|
||||
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
|
||||
- ``weight_type``: ``beta08``
|
||||
- ``fuse_guidance_scale``: 3.0
|
||||
- Training steps: 6000
|
||||
|
||||
### Stage 2 — on-policy
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
|
||||
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
|
||||
```
|
||||
|
||||
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
|
||||
|
||||
**Key configuration**:
|
||||
|
||||
- Global batch size: 8 (8 GPUs × 1 per-GPU)
|
||||
- Learning rate: 2e-6
|
||||
- Flow shift: 5.0
|
||||
- ``student_sample_steps``: 4
|
||||
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
|
||||
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
|
||||
- ``real_score_guidance_scale``: 3.0
|
||||
- ``generator_update_interval``: 5 (DMD2 alternation)
|
||||
- Training steps: 4000
|
||||
|
||||
## 🔌 Loading published AnyFlow checkpoints
|
||||
|
||||
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
|
||||
|
||||
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
|
||||
|
||||
## 🧭 Note on ``fuse_guidance_scale``
|
||||
|
||||
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
|
||||
|
||||
```
|
||||
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
|
||||
```
|
||||
|
||||
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
|
||||
|
||||
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
|
||||
@@ -74,6 +74,23 @@ uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
### Flash Attention 4 (opt-in)
|
||||
|
||||
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
|
||||
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
|
||||
runtime on some GPU/shape combinations. To use FA4, install the pinned
|
||||
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
|
||||
|
||||
```bash
|
||||
export FASTVIDEO_FA4=1
|
||||
```
|
||||
|
||||
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
|
||||
cannot serve there: grad-enabled (training) attention (FA4's backward requires
|
||||
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
|
||||
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
|
||||
FastVideo fails loudly instead of silently falling back.
|
||||
|
||||
### FP4 Flash Attention 4 (Blackwell only)
|
||||
|
||||
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
|
||||
|
||||
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
# Training Trackers
|
||||
|
||||
FastVideo can send training metrics and validation media to Weights & Biases
|
||||
or SwanLab. Tracking runs only on global rank 0, and local tracker files are
|
||||
stored under `<output_dir>/tracker`.
|
||||
|
||||
## Supported Trackers
|
||||
|
||||
| Value | Backend | Installation |
|
||||
|-------|---------|--------------|
|
||||
| `wandb` | Weights & Biases | Included with FastVideo |
|
||||
| `swanlab` | SwanLab | Install the optional `swanlab` dependency |
|
||||
| `none` | Disable external tracking | No additional package |
|
||||
|
||||
You can enable more than one backend, for example `trackers: [wandb, swanlab]`.
|
||||
Metrics and validation media are converted to the artifact type required by
|
||||
each backend.
|
||||
|
||||
## Install SwanLab
|
||||
|
||||
For a published FastVideo installation, install the SwanLab extra:
|
||||
|
||||
```bash
|
||||
uv pip install "fastvideo[swanlab]"
|
||||
```
|
||||
|
||||
For an editable source checkout, include the same extra during installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[swanlab]"
|
||||
```
|
||||
|
||||
If FastVideo is already installed, you can install the compatible SDK directly:
|
||||
|
||||
```bash
|
||||
uv pip install "swanlab>=0.6.7"
|
||||
```
|
||||
|
||||
Authenticate once before starting a training run:
|
||||
|
||||
```bash
|
||||
swanlab login
|
||||
```
|
||||
|
||||
See the [SwanLab login documentation](https://docs.swanlab.cn/en/api/cli-swanlab-login.html)
|
||||
for non-interactive and self-hosted setups.
|
||||
|
||||
## Configure Tracking
|
||||
|
||||
Select SwanLab in the YAML config used by the modular training framework:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
checkpoint:
|
||||
output_dir: outputs/my_run
|
||||
tracker:
|
||||
trackers: [swanlab]
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
To log to both supported services:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
trackers: [wandb, swanlab]
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
An empty or omitted `trackers` list selects W&B when `project_name` is set.
|
||||
Use an explicit `none` entry to disable external tracking:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
trackers: [none]
|
||||
```
|
||||
|
||||
## Validation Videos
|
||||
|
||||
SwanLab currently accepts GIF video artifacts. FastVideo converts validation
|
||||
MP4 files and in-memory video arrays to GIF automatically before logging them.
|
||||
For video files, FastVideo uses the sampling frame rate supplied by the caller,
|
||||
or the source file's frame rate when no value is supplied. In-memory arrays use
|
||||
the frame rate supplied by the caller. Both forms fall back to 16 FPS when no
|
||||
frame rate is available.
|
||||
|
||||
For details about configuring validation callbacks, see
|
||||
[Training Infrastructure](train_infra.md#callbacks-pluggable-hooks).
|
||||
@@ -161,6 +161,21 @@ training:
|
||||
decay_interval_steps: 0
|
||||
```
|
||||
|
||||
`training.data.data_path` can also mix multiple preprocessed datasets by using a mapping from dataset path to repeat count:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
data/multi3d_games: 2
|
||||
```
|
||||
|
||||
The repeat count duplicates that dataset's parquet file list before shuffling/sampling, so the example above trains with roughly twice as much `multi3d_games` exposure as `zeldam2-clean`. Paths are just suggested locations; use any local path that contains a FastVideo preprocessed parquet dataset.
|
||||
|
||||
See [Training Trackers](trackers.md) to configure Weights & Biases or SwanLab,
|
||||
including SwanLab installation and authentication.
|
||||
|
||||
### `callbacks` — Pluggable hooks
|
||||
|
||||
Callbacks run at specific points in the training loop (before/after optimizer
|
||||
@@ -323,6 +338,40 @@ Self-Forcing inherits all DMD2 parameters, plus:
|
||||
| `enable_gradient_in_rollout` | `true` | Enable backprop through rollout |
|
||||
| `start_gradient_frame` | `0` | Frame index where gradients begin |
|
||||
|
||||
### Streaming Long Tuning
|
||||
|
||||
`StreamingLongTuningMethod` extends Self-Forcing for LongLive-style rollouts. It
|
||||
keeps a streaming state, generates overlapping chunks, and trains only the new
|
||||
frames while preserving context from earlier chunks.
|
||||
|
||||
For the MatrixGame2/Zelda world-model example, self-forcing and long tuning are
|
||||
separate runs: first train or load the 1k-step self-forcing checkpoint using
|
||||
`examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml`,
|
||||
then run
|
||||
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
|
||||
from that checkpoint for the 3k-step streaming long-tuning stage.
|
||||
|
||||
```yaml
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
streaming_reencode_overlap_anchor: true
|
||||
streaming_anchor_inject_k: 1
|
||||
streaming_require_full_blocks: true
|
||||
multi_phased_distill_schedule:
|
||||
- stage: streaming_long
|
||||
start_step: 0
|
||||
end_step: 3000
|
||||
num_latent_t: 39
|
||||
streaming_training: true
|
||||
```
|
||||
|
||||
See
|
||||
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
|
||||
for a complete MatrixGame2/Zelda configuration.
|
||||
|
||||
---
|
||||
|
||||
## Callbacks
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
|
||||
|
||||
def _env_int(name: str, default: int) -> int:
|
||||
return int(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def _env_float(name: str, default: float) -> float:
|
||||
return float(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
|
||||
prompt = os.getenv(
|
||||
"DREAMX_WORLD_PROMPT",
|
||||
"A cinematic first-person drive through a futuristic coastal city at "
|
||||
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
|
||||
)
|
||||
image_path = os.getenv(
|
||||
"DREAMX_WORLD_IMAGE_PATH",
|
||||
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
|
||||
)
|
||||
|
||||
kwargs = {
|
||||
"output_path": OUTPUT_PATH,
|
||||
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
|
||||
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
|
||||
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
|
||||
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
|
||||
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
|
||||
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
|
||||
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
|
||||
"action_speed_list": [
|
||||
float(value)
|
||||
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
|
||||
],
|
||||
}
|
||||
if image_path:
|
||||
kwargs["image_path"] = image_path
|
||||
|
||||
try:
|
||||
generator.generate_video(prompt, **kwargs)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,140 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import os
|
||||
import re
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
(
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _safe_filename(text: str, max_len: int = 100) -> str:
|
||||
"""Make a stable, filesystem-friendly filename base."""
|
||||
s = text[:max_len].strip()
|
||||
s = s.replace(os.sep, "_")
|
||||
if os.altsep:
|
||||
s = s.replace(os.altsep, "_")
|
||||
s = re.sub(r"\s+", " ", s)
|
||||
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
|
||||
s = s.strip(" .")
|
||||
return s or "prompt"
|
||||
|
||||
|
||||
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
|
||||
if not os.path.isdir(out_dir):
|
||||
return
|
||||
|
||||
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
|
||||
for fn in os.listdir(out_dir):
|
||||
if pattern.match(fn):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.remove(os.path.join(out_dir, fn))
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--model-path",
|
||||
default="official_weights/FLUX.1-dev",
|
||||
help="Local Diffusers checkpoint dir or HF repo id.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
default="outputs/flux_dev/samples",
|
||||
help="Directory for saved PNG outputs.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--prompt",
|
||||
action="append",
|
||||
default=None,
|
||||
help="Prompt. Repeat for multiple images.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
|
||||
)
|
||||
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
|
||||
p.add_argument("--height", type=int, default=1024, help="Output height.")
|
||||
p.add_argument("--width", type=int, default=1024, help="Output width.")
|
||||
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
|
||||
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
|
||||
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
|
||||
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": args.num_gpus,
|
||||
"workload_type": "t2i",
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"image_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": False,
|
||||
"use_fsdp_inference": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=args.model_path,
|
||||
**init_kwargs,
|
||||
)
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
filename_base = (
|
||||
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
)
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
generation_kwargs = {
|
||||
"output_path": output_path,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": args.steps,
|
||||
"guidance_scale": args.guidance,
|
||||
"use_embedded_guidance": True,
|
||||
"true_cfg_scale": 1.0,
|
||||
"seed": seed,
|
||||
"save_video": True,
|
||||
}
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
print(f"[flux] done. outputs written to: {args.out_dir}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,107 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
|
||||
text-to-image generation command, saved as a PNG."
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/landscape.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default=("A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style."),
|
||||
help="Text prompt.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
|
||||
|
||||
IMAGE_PATH = "assets/girl.png"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A woman stands up and walks away"
|
||||
)
|
||||
_ = generator.generate_video(
|
||||
prompt,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=1024,
|
||||
width=1024,
|
||||
num_frames=121,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,120 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
|
||||
want a minimal edit command (text + image -> edited image), saved as a PNG."
|
||||
|
||||
GLM-Image is a single unified pipeline: passing a condition image switches it
|
||||
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
|
||||
write pass), so the generator config is identical to `basic_glm_image.py` — the
|
||||
`inputs.pil_image` on the request is what selects the edit mode.
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image",
|
||||
default="assets/images/couple.jpg",
|
||||
help="Condition image to edit.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/edited.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="Change the background to a snowy mountain landscape at golden hour.",
|
||||
help="Edit instruction.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
condition = Image.open(args.image).convert("RGB")
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
# The pipeline is registered as t2i; passing inputs.pil_image below switches
|
||||
# it to the edit path.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
inputs=InputConfig(pil_image=condition),
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,69 @@
|
||||
# LTX-2.3 distilled inference configs
|
||||
|
||||
Ready-to-run `fastvideo generate` run configs for the LTX-2.3
|
||||
distilled model (`FastVideo/LTX-2.3-Distilled-Diffusers`), covering both
|
||||
workloads (t2v / i2v), both two-stage step schedules (`5+2`, `8+3` = denoise
|
||||
+ refine), and four resolutions.
|
||||
|
||||
```bash
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
|
||||
```
|
||||
|
||||
Each config is self-contained (no preset registry needed): the two-stage
|
||||
refine is wired via `generator.pipeline.preset_overrides.refine`, and the
|
||||
base sampling knobs live under `request.sampling`. The refine upsampler
|
||||
auto-resolves from the model's `spatial_upscaler`.
|
||||
|
||||
## Configs
|
||||
|
||||
| workload | schedule | resolution (HxW) | file |
|
||||
|---|---|---|---|
|
||||
| t2v | 5+2 | 1280x832 | `t2v_5s2_1280x832.yaml` |
|
||||
| t2v | 5+2 | 1024x1536 | `t2v_5s2_1024x1536.yaml` |
|
||||
| t2v | 5+2 | 768x1280 | `t2v_5s2_768x1280.yaml` |
|
||||
| t2v | 5+2 | 512x768 | `t2v_5s2_512x768.yaml` |
|
||||
| t2v | 8+3 | 1280x832 | `t2v_8s3_1280x832.yaml` |
|
||||
| t2v | 8+3 | 1024x1536 | `t2v_8s3_1024x1536.yaml` |
|
||||
| t2v | 8+3 | 768x1280 | `t2v_8s3_768x1280.yaml` |
|
||||
| t2v | 8+3 | 512x768 | `t2v_8s3_512x768.yaml` |
|
||||
| i2v | 5+2 | 1280x832 | `i2v_5s2_1280x832.yaml` |
|
||||
| i2v | 5+2 | 1024x1536 | `i2v_5s2_1024x1536.yaml` |
|
||||
| i2v | 5+2 | 768x1280 | `i2v_5s2_768x1280.yaml` |
|
||||
| i2v | 5+2 | 512x768 | `i2v_5s2_512x768.yaml` |
|
||||
| i2v | 8+3 | 1280x832 | `i2v_8s3_1280x832.yaml` |
|
||||
| i2v | 8+3 | 1024x1536 | `i2v_8s3_1024x1536.yaml` |
|
||||
| i2v | 8+3 | 768x1280 | `i2v_8s3_768x1280.yaml` |
|
||||
| i2v | 8+3 | 512x768 | `i2v_8s3_512x768.yaml` |
|
||||
|
||||
## Overriding without editing a file
|
||||
|
||||
Dotted overrides (prefixes `generator.` / `request.`) let you tweak any field:
|
||||
|
||||
```bash
|
||||
# swap prompt
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml \
|
||||
--request.prompt "a red fox running through fresh snow"
|
||||
|
||||
# change output path / gpu count
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml \
|
||||
--request.output.output_path outputs/preview.mp4 \
|
||||
--generator.engine.num_gpus 4
|
||||
```
|
||||
|
||||
## i2v
|
||||
|
||||
The `i2v_*` configs take a first-frame image via
|
||||
`request.extensions.ltx2_images` (`[[path, frame_offset, weight]]`). Edit the
|
||||
path in the file, or override it:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml \
|
||||
--request.extensions.ltx2_images '[["/data/portrait.jpg", 0, 1.0]]'
|
||||
```
|
||||
|
||||
## Schedules
|
||||
|
||||
`5+2` is the fast preview schedule; `8+3` is the higher-quality distilled
|
||||
recipe. Refine (`preset_overrides.refine.num_inference_steps`) only accepts 2
|
||||
or 3 steps.
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,78 @@
|
||||
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
|
||||
# the student matches an EMA copy of itself at the next timestep, all under
|
||||
# clean-history teacher forcing.
|
||||
#
|
||||
# All three roles initialize from the SAME checkpoint (the teacher-forcing
|
||||
# AR-diffusion model). Point init_from at that checkpoint for a real run.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 48
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.99
|
||||
ema_start_step: 200
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_cd
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_cd_shift5
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,105 @@
|
||||
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
|
||||
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
|
||||
# rollout from pure noise. Teacher provides the real score, critic
|
||||
# learns the fake score; both inherited from DMD2Method.
|
||||
#
|
||||
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
|
||||
# or with the NVIDIA-released checkpoint
|
||||
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
|
||||
# the paper weights (the delta_embedder rename is handled by the
|
||||
# param_names_mapping in WanVideoArchConfig).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: <PATH_TO_PRETRAIN_CKPT>
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 3.0
|
||||
dmd_denoising_steps: [999, 937, 833, 624]
|
||||
warp_denoising_step: false
|
||||
|
||||
# AnyFlow rollout knobs.
|
||||
student_sample_steps: 4
|
||||
use_mean_velocity: true
|
||||
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
|
||||
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
|
||||
|
||||
# Critic optimizer (DMD2 inherited).
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
attn_kind: vsa
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_onpolicy
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_onpolicy
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -0,0 +1,83 @@
|
||||
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
|
||||
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
|
||||
# can be sampled at arbitrary NFE in the on-policy stage.
|
||||
#
|
||||
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
|
||||
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
|
||||
diffusion_ratio: 0.5
|
||||
consistency_ratio: 0.25
|
||||
epsilon: 5 # finite-difference step in absolute train-timestep units
|
||||
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
|
||||
fuse_guidance_scale: 3.0
|
||||
# shift is taken from pipeline.flow_shift below.
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 4
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_pretrain
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_pretrain
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
# Enable AnyFlow dual-timestep conditioning. The student loads from
|
||||
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
|
||||
# get initialized identically to time_embedder via deep-copy in
|
||||
# WanTimeTextImageEmbedding.__init__.
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -100,9 +100,6 @@ callbacks:
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 81
|
||||
# Validation/inference uses standard CFG in both clean and Self-Forcing,
|
||||
# so this directly matches Self-Forcing guidance_scale=3.0.
|
||||
guidance_scale: 3.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model with a block size of 1 frame
|
||||
# - Training: each frame gets its own independent noise level (frame-wise
|
||||
# diffusion forcing), versus the chunk-wise variant that shares one noise
|
||||
# level across num_frames_per_block frames.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_dfsft_framewise
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,72 @@
|
||||
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model
|
||||
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
|
||||
# denoises the current block while attending to *clean* history (clean_x),
|
||||
# not its own noisy rollout.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_tfsft
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,32 +1,119 @@
|
||||
# World-Model: Matrix-Game 2.0 I2V
|
||||
|
||||
Three training scenarios for the Matrix-Game 2.0 I2V world model on the
|
||||
new YAML-driven trainer (`fastvideo/train/entrypoint/train.py`).
|
||||
Training scenarios for the Matrix-Game 2.0 I2V world model on Solaris (Minecraft)
|
||||
data and Zelda data, using the new YAML-driven trainer
|
||||
(`fastvideo/train/entrypoint/train.py`).
|
||||
|
||||
## Solaris Configs
|
||||
|
||||
| Config | Method | Student | Notes |
|
||||
|---|---|---|---|
|
||||
| `finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
|
||||
| `dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
|
||||
| `self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
|
||||
| `solaris/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
|
||||
| `solaris/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
|
||||
| `solaris/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Matrix-Game 2.0 DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
|
||||
|
||||
## Zelda Configs
|
||||
|
||||
| Config | Method | Student | Notes |
|
||||
|---|---|---|---|
|
||||
| `zelda/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Zelda bidirectional I2V finetuning from `FastVideo/Matrix-Game-2.0-Base-Diffusers`. Uses 33-frame clips and Zelda validation with action overlays. |
|
||||
| `zelda/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Zelda causal Diffusion-Forcing SFT from `mignonjia/mg_bidirectional_zelda`. Uses the same Zelda data, resolution, optimizer, and validation defaults as the Zelda finetune config. |
|
||||
| `zelda/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Zelda DMD/Self-Forcing distillation; student init = `mignonjia/mg_causal_zelda`, teacher = bidirectional (`mignonjia/mg_bidirectional_zelda`), critic = bidirectional. |
|
||||
| `zelda/streaming_long_tuning_causal_i2v.yaml` | `StreamingLongTuningMethod` | `MatrixGame2CausalModel` | LongLive-style streaming long tuning from the 1k-step Zelda self-forcing checkpoint. |
|
||||
|
||||
Zelda world-model distillation is a two-run workflow: first run
|
||||
`zelda/self_forcing_causal_i2v.yaml` to train or load the 1k-step
|
||||
self-forcing checkpoint (`mignonjia/mg_sf_distilled_zelda_1k_steps`), then run
|
||||
`zelda/streaming_long_tuning_causal_i2v.yaml` for the 3k-step streaming
|
||||
long-tuning stage. The long-tuning YAML starts from that 1k-step checkpoint; it
|
||||
does not run the short self-forcing stage inside the same config.
|
||||
|
||||
## Zelda Training Data
|
||||
|
||||
The Zelda training configs use `data/zeldam2-clean` as a suggested local path.
|
||||
Download the dataset from Hugging Face before running those configs:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id mignonjia/zeldam2-clean \
|
||||
--local_dir data/zeldam2-clean \
|
||||
--repo_type dataset
|
||||
```
|
||||
|
||||
You can store the dataset elsewhere; update `training.data.data_path` in the
|
||||
YAML to point at that location.
|
||||
|
||||
## Multi3D Training Data
|
||||
|
||||
`zelda/finetune_i2v.yaml` includes an optional, commented-out Multi3D entry.
|
||||
Enable it only when you want to mix Zelda with multi-game data from
|
||||
`data/multi3d_games`. You can store this dataset anywhere; before enabling it,
|
||||
update the matching commented `training.data.data_path` key in the YAML to the
|
||||
correct location.
|
||||
|
||||
To mix datasets in a training YAML, set `training.data.data_path` to a
|
||||
path-to-repeat-count mapping. For example, `zelda/finetune_i2v.yaml` can use
|
||||
`data/zeldam2-clean: 1` and `# data/multi3d_games: 10`; uncommenting the
|
||||
Multi3D entry repeats the multi-game parquet list ten times before training
|
||||
samples are shuffled.
|
||||
|
||||
## World Model Validation Data
|
||||
|
||||
The Zelda validation configs expect a small public validation bundle under
|
||||
`data/zelda_validation_data`.
|
||||
|
||||
Download it from Hugging Face before running the Zelda scenarios:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id mignonjia/zelda_validation_data \
|
||||
--local_dir data/zelda_validation_data \
|
||||
--repo_type dataset
|
||||
```
|
||||
|
||||
The bundle contains `validation_zelda.json`, `images/`, and `actions/`.
|
||||
The Zelda configs point
|
||||
`callbacks.validation.dataset_file` at
|
||||
`data/zelda_validation_data/validation_zelda.json`.
|
||||
|
||||
## Usage
|
||||
|
||||
### Solaris
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/finetune_i2v.yaml
|
||||
examples/train/scenario/worldmodel/solaris/finetune_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml
|
||||
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/self_forcing_causal_i2v.yaml
|
||||
examples/train/scenario/worldmodel/solaris/self_forcing_causal_i2v.yaml
|
||||
```
|
||||
|
||||
### Zelda
|
||||
|
||||
```bash
|
||||
# Finetuning / DFSFT
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/finetune_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/dfsft_causal_i2v.yaml
|
||||
|
||||
# Distillation / long tuning
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml
|
||||
```
|
||||
|
||||
Override any field on the command line:
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml \
|
||||
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml \
|
||||
--training.distributed.num_gpus 8 \
|
||||
--training.optimizer.learning_rate 1e-5
|
||||
```
|
||||
|
||||
+1
-1
@@ -97,4 +97,4 @@ callbacks:
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,94 @@
|
||||
# Diffusion-Forcing SFT: Zelda world model I2V Causal
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-5
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_causal_dfsft
|
||||
training_state_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
entity: hapo-exp
|
||||
project_name: mg_1.3b_zelda
|
||||
run_name: zelda_causal_dfsft
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 200
|
||||
sampling_steps: [40]
|
||||
sampling_timesteps: [1000, 975, 950, 925, 900, 875, 850, 825, 800, 775,
|
||||
750, 725, 700, 675, 650, 625, 600, 575, 550, 525,
|
||||
500, 475, 450, 425, 400, 375, 350, 325, 300, 275,
|
||||
250, 225, 200, 175, 150, 125, 100, 75, 50, 25]
|
||||
num_frames: 33
|
||||
overlay_actions: true
|
||||
guidance_scale: 6.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -0,0 +1,88 @@
|
||||
# Matrix-Game 2.0 Zelda + multi-game I2V finetune.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: FastVideo/Matrix-Game-2.0-Base-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
# data/multi3d_games: 10
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0 # unused for MatrixGame2 I2V; no text_embedding CFG dropout
|
||||
seed: 42
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-5
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_with_mg_init
|
||||
training_state_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: mg_1.3b_zelda
|
||||
run_name: zelda_with_mg_init
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_i2v_pipeline.MatrixGame2I2VPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 200
|
||||
sampling_steps: [40]
|
||||
num_frames: 33
|
||||
overlay_actions: true
|
||||
guidance_scale: 6.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,120 @@
|
||||
# Self-forcing distillation: Zelda world model I2V Causal
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_causal_zelda
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
same_step_across_blocks: true
|
||||
last_step_only: false
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 3.0e-7
|
||||
fake_score_betas: [0.9, 0.95]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/zeldam2-clean
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1001
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 3.0e-6
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/zelda_causal_self_forcing
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: mg2_self_forcing_9_latents
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
ema:
|
||||
decay: 0.99
|
||||
start_iter: 200
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 100
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 153
|
||||
overlay_actions: true
|
||||
keyboard_value_scale: 1.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -0,0 +1,140 @@
|
||||
# MatrixGame2 I2V LongLive-style streaming distillation: Zelda world model I2V Causal
|
||||
# Student init from self forcing after 1k steps
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_sf_distilled_zelda_1k_steps
|
||||
trainable: true
|
||||
# transformer_override_safetensor: outputs/matrixgame_dmd/checkpoints/mg_zelda_sf_m2/checkpoint-1000_weight_only/ema/generator_ema.safetensors
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
same_step_across_blocks: true
|
||||
last_step_only: false
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
streaming_training: true
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
streaming_reencode_overlap_anchor: true
|
||||
streaming_anchor_inject_k: 1
|
||||
streaming_require_full_blocks: true
|
||||
|
||||
multi_phased_distill_schedule:
|
||||
- stage: streaming_long
|
||||
start_step: 0
|
||||
end_step: 3000
|
||||
num_latent_t: 39
|
||||
streaming_training: true
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 3.0e-7
|
||||
fake_score_betas: [0.9, 0.95]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/zeldam2-clean
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1001
|
||||
num_latent_t: 39
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 153
|
||||
|
||||
optimizer:
|
||||
learning_rate: 3.0e-6
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/zelda_causal_long_tuning
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: mg2_39only_streaming_long
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
ema:
|
||||
decay: 0.99
|
||||
start_iter: 200
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 100
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 153
|
||||
overlay_actions: true
|
||||
keyboard_value_scale: 1.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -173,6 +185,14 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,6 +2,42 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,6 +46,23 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxSamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 1024
|
||||
width: int = 1024
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 3.5
|
||||
use_embedded_guidance: bool = True
|
||||
true_cfg_scale: float = 1.0
|
||||
@@ -90,6 +90,10 @@ class SamplingParam:
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
|
||||
use_embedded_guidance: bool = False
|
||||
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
|
||||
true_cfg_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
@@ -325,6 +329,18 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-embedded-guidance",
|
||||
action="store_true",
|
||||
default=SamplingParam.use_embedded_guidance,
|
||||
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--true-cfg-scale",
|
||||
type=float,
|
||||
default=SamplingParam.true_cfg_scale,
|
||||
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
|
||||
@@ -150,6 +150,7 @@ class SamplingConfig:
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
use_embedded_guidance: bool | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@@ -5,99 +5,10 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
|
||||
fa_version = "4"
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer have a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
|
||||
# breaks the graph at the call site (observed: wanvideo.py self-attn,
|
||||
# once per layer every step), which fragments the compiled region and
|
||||
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
|
||||
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
# opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
# (correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
|
||||
if fa_version in ("2", "3"):
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
|
||||
# call shape used by FlashAttentionImpl.forward's default branch
|
||||
# (see `flash_attn_func_compilable(...)` call site below). The
|
||||
# masked/no-pad and varlen / cross-attn paths use different
|
||||
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
|
||||
# are intentionally out of scope for this PR — wrapping them is a
|
||||
# natural follow-up. The wrapper's signature is the contract: any
|
||||
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
|
||||
# return_attn_probs, ...) raises TypeError at the call site, so
|
||||
# silent loss of kwargs is not a failure mode.
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
|
||||
# same dtype/device as q (head dim taken from v).
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
from fastvideo.attention.utils.flash_attn_default import (
|
||||
fa_version,
|
||||
flash_attn_func_compilable,
|
||||
)
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
@@ -108,8 +19,6 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_WARNED_NON_FA_DTYPE = False
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
@@ -271,12 +180,8 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
|
||||
orig_dtype = query.dtype
|
||||
if orig_dtype not in (torch.float16, torch.bfloat16):
|
||||
global _WARNED_NON_FA_DTYPE
|
||||
if not _WARNED_NON_FA_DTYPE:
|
||||
_WARNED_NON_FA_DTYPE = True
|
||||
logger.warning(
|
||||
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
|
||||
"kernel and restoring on output.", orig_dtype)
|
||||
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
|
||||
f"bfloat16 for the kernel and restoring on output.")
|
||||
query = query.to(torch.bfloat16)
|
||||
key = key.to(torch.bfloat16)
|
||||
value = value.to(torch.bfloat16)
|
||||
@@ -293,9 +198,17 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
|
||||
# Route through the *_compilable wrappers so dynamo sees one
|
||||
# traceable node for each masked entry point (the unpad/pad
|
||||
# bookkeeping runs eager inside the custom op). On FA2 these
|
||||
# wrappers go through ops with full register_autograd, so
|
||||
# training also backprops through the op (no graph break on
|
||||
# the training path); on FA3/FA4 they carve out to the
|
||||
# autograd.Function for grad-enabled calls — see
|
||||
# fastvideo/attention/utils/flash_attn_no_pad.py.
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad,
|
||||
flash_attn_no_pad_compilable as flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad_compilable as flash_attn_varlen_qk_no_pad,
|
||||
)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
@@ -322,7 +235,11 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
)
|
||||
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
attn_mask_padded = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
|
||||
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, attn_mask.shape[-1]).to(device=query.device)
|
||||
if key_padding_mask.shape[-1] > qkv.shape[1]:
|
||||
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
|
||||
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
|
||||
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
|
||||
elif self.nvfp4_fa4:
|
||||
output = self._forward_nvfp4(query, key, value)
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NABLA block-sparse flex-attention backend (Kandinsky5 "nabla" checkpoints).
|
||||
|
||||
The block mask is data-dependent: nablaT_v2 mean-pools 64-token blocks of the
|
||||
fractal-ordered sequence, thresholds the softmaxed block map, and ORs it with a
|
||||
precomputed spatio-temporal-window (STA) mask carried on the attention
|
||||
metadata. The mask spans the full sequence, so this backend does not support
|
||||
sequence parallelism — use it via LocalAttention only.
|
||||
"""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
from torch.nn.attention.flex_attention import BlockMask, flex_attention
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
|
||||
CAN_USE_FLEX_ATTN = True
|
||||
except ImportError:
|
||||
CAN_USE_FLEX_ATTN = False
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
|
||||
|
||||
def nablaT_v2(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
sta: torch.Tensor,
|
||||
thr: float = 0.9,
|
||||
) -> "BlockMask":
|
||||
q = q.transpose(1, 2).contiguous()
|
||||
k = k.transpose(1, 2).contiguous()
|
||||
|
||||
# Map estimation
|
||||
B, h, S, D = q.shape
|
||||
s1 = S // 64
|
||||
qa = q.reshape(B, h, s1, 64, D).mean(-2)
|
||||
ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1)
|
||||
map = qa @ ka
|
||||
|
||||
map = torch.softmax(map / math.sqrt(D), dim=-1)
|
||||
# Map binarization
|
||||
vals, inds = map.sort(-1)
|
||||
cvals = vals.cumsum_(-1)
|
||||
mask = (cvals >= 1 - thr).int()
|
||||
mask = mask.gather(-1, inds.argsort(-1))
|
||||
|
||||
mask = torch.logical_or(mask, sta)
|
||||
|
||||
# BlockMask creation
|
||||
kv_nb = mask.sum(-1).to(torch.int32)
|
||||
kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32)
|
||||
return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None)
|
||||
|
||||
|
||||
class NablaAttentionBackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "NABLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["NablaAttentionImpl"]:
|
||||
return NablaAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["NablaAttentionMetadata"]:
|
||||
return NablaAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["NablaAttentionMetadataBuilder"]:
|
||||
return NablaAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class NablaAttentionMetadata(AttentionMetadata):
|
||||
# Block-level STA window mask [1, 1, S/64, S/64], precomputed once per run.
|
||||
sta_mask: torch.Tensor = None # type: ignore[assignment]
|
||||
# Cumulative-probability threshold for block-map binarization.
|
||||
P: float = 0.9
|
||||
visual_shape: tuple[int, int, int] = (0, 0, 0)
|
||||
|
||||
|
||||
class NablaAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
sta_mask: torch.Tensor,
|
||||
P: float,
|
||||
visual_shape: tuple[int, int, int],
|
||||
**kwargs: Any,
|
||||
) -> NablaAttentionMetadata:
|
||||
return NablaAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
sta_mask=sta_mask,
|
||||
P=P,
|
||||
visual_shape=visual_shape,
|
||||
)
|
||||
|
||||
|
||||
class NablaAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
causal: bool = False,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
if not CAN_USE_FLEX_ATTN:
|
||||
raise RuntimeError("NABLA attention requires torch.nn.attention.flex_attention, "
|
||||
"which is unavailable in this PyTorch build.")
|
||||
if causal:
|
||||
raise ValueError("NABLA attention does not support causal masking.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: NablaAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
# q/k/v: [B, S, heads, head_dim], fractal-ordered by the model; S % 64 == 0.
|
||||
block_mask = nablaT_v2(query, key, attn_metadata.sta_mask, thr=attn_metadata.P)
|
||||
return flex_attention(
|
||||
query=query.transpose(1, 2),
|
||||
key=key.transpose(1, 2),
|
||||
value=value.transpose(1, 2),
|
||||
block_mask=block_mask,
|
||||
).transpose(1, 2)
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
from torch.nn import functional as F
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -14,12 +15,18 @@ class SDPABackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
def get_supported_head_sizes() -> list[int] | None:
|
||||
# torch.nn.functional.scaled_dot_product_attention is head-size
|
||||
# agnostic: the math backend handles any size and the fused kernels
|
||||
# fall back internally when they cannot. The previous list
|
||||
# ([32, 64, ..., 256]) was copied from FlashAttentionBackend in the
|
||||
# v1 refactor (#270) and wrongly excluded sizes such as 80 (e.g.
|
||||
# CLIP vision encoders). None means "no restriction".
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SDPA"
|
||||
return "TORCH_SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SDPAImpl"]:
|
||||
@@ -49,9 +56,51 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
# Store the mask exactly as passed. The metadata is cross-backend:
|
||||
# call sites (HYWorld, HunyuanVideo15) build SDPAMetadata while the
|
||||
# layer's selector may pick FLASH_ATTN, and the shared convention for
|
||||
# padding masks is the tokenizer-style 2D [batch, key_len]. Any
|
||||
# reshaping for torch.sdpa happens inside the SDPA impl
|
||||
# (_normalize_attn_mask_for_sdpa).
|
||||
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
|
||||
|
||||
def _normalize_attn_mask_for_sdpa(
|
||||
attn_mask: torch.Tensor | None,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
) -> torch.Tensor | None:
|
||||
if attn_mask is None:
|
||||
return None
|
||||
|
||||
attn_mask = attn_mask.to(device=query.device)
|
||||
# F.scaled_dot_product_attention only accepts bool or float masks;
|
||||
# tokenizers commonly produce int64 0/1 padding masks.
|
||||
if attn_mask.dtype != torch.bool and not attn_mask.dtype.is_floating_point:
|
||||
attn_mask = attn_mask != 0
|
||||
|
||||
key_len = key.shape[-2]
|
||||
if attn_mask.shape[-1] > key_len:
|
||||
raise ValueError("Invalid attention mask length for SDPA: "
|
||||
f"expected at most {key_len}, got {attn_mask.shape[-1]}")
|
||||
if attn_mask.shape[-1] < key_len:
|
||||
# Front-pad as "attend": double-stream layouts (HYWorld) prepend
|
||||
# non-text tokens the tokenizer mask does not cover.
|
||||
valid_value = True if attn_mask.dtype == torch.bool else 0.0
|
||||
attn_mask = F.pad(attn_mask, (key_len - attn_mask.shape[-1], 0), value=valid_value)
|
||||
|
||||
if attn_mask.dim() == 2:
|
||||
# In-tree producers pass 2D [batch, key_len] padding masks; lift to a
|
||||
# broadcastable [batch, 1, 1, key_len] here so torch.sdpa does not
|
||||
# reinterpret 2D as its documented [query_len, key_len] broadcast.
|
||||
return attn_mask[:, None, None, :]
|
||||
if attn_mask.dim() == 3:
|
||||
return attn_mask[:, None, :, :]
|
||||
if attn_mask.dim() == 4:
|
||||
return attn_mask
|
||||
raise ValueError(f"Unsupported attention mask shape for SDPA: {attn_mask.shape}")
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -82,6 +131,7 @@ class SDPAImpl(AttentionImpl):
|
||||
|
||||
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
|
||||
and hasattr(attn_metadata, "attn_mask")) else None
|
||||
attn_mask = _normalize_attn_mask_for_sdpa(attn_mask, query, key)
|
||||
attn_kwargs = {
|
||||
"attn_mask": attn_mask,
|
||||
"dropout_p": self.dropout,
|
||||
|
||||
@@ -252,6 +252,7 @@ class LocalAttention(nn.Module):
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -262,7 +263,10 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
|
||||
attn_backend = get_attn_backend(head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
default_backend=default_backend)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
|
||||
@@ -84,8 +84,9 @@ def get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
|
||||
|
||||
|
||||
@cache
|
||||
@@ -94,6 +95,7 @@ def _cached_get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -112,6 +114,12 @@ def _cached_get_attn_backend(
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
# Layer-level default (e.g. a checkpoint that requires a specific sparse
|
||||
# backend). Lower precedence than the global force and the env var, so
|
||||
# users can still override it.
|
||||
if selected_backend is None and default_backend is not None:
|
||||
selected_backend = default_backend
|
||||
|
||||
# get device-specific attn_backend
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ import functools
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -16,7 +15,8 @@ if torch.cuda.is_available():
|
||||
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
|
||||
except ImportError:
|
||||
# flash_attn.cute (FA4) is simply not installed -- expected on builds
|
||||
# without it; callers fall back to FA3/FA2 quietly.
|
||||
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
|
||||
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
|
||||
raise
|
||||
except Exception as e:
|
||||
# flash_attn.cute IS installed but failed to import -- almost always an
|
||||
@@ -24,23 +24,57 @@ if torch.cuda.is_available():
|
||||
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
|
||||
# ImportError). This is fixable by pinning a compatible
|
||||
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
|
||||
# callers fall back to FA3/FA2 instead of crashing worker init.
|
||||
# callers can handle it uniformly.
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r); "
|
||||
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
|
||||
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
|
||||
"restore FA4.", e)
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r). "
|
||||
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
|
||||
"compatible nvidia-cutlass-dsl to restore FA4.", e)
|
||||
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
|
||||
else:
|
||||
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
|
||||
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
|
||||
|
||||
try:
|
||||
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
|
||||
# Optional so FA4-only installs can still import this module.
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
except ImportError:
|
||||
_flash_attn_2_func = None
|
||||
_flash_attn_2_varlen_func = None
|
||||
|
||||
|
||||
def _check_dropout(dropout_p: float) -> None:
|
||||
if dropout_p != 0.0:
|
||||
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _sm90_or_newer() -> bool:
|
||||
return current_platform.has_device_capability(90)
|
||||
|
||||
|
||||
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
|
||||
if _sm90_or_newer():
|
||||
return False
|
||||
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
|
||||
# capability gate, not a runtime fallback):
|
||||
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
|
||||
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
|
||||
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
|
||||
if q.shape[-2] != k.shape[-2]:
|
||||
return True
|
||||
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
|
||||
|
||||
|
||||
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
|
||||
if fa2_func is None:
|
||||
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
|
||||
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
|
||||
"not installed.")
|
||||
return fa2_func
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_cute_forward",
|
||||
mutates_args=(),
|
||||
@@ -243,70 +277,6 @@ torch.library.register_autograd(
|
||||
)
|
||||
|
||||
|
||||
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
|
||||
# fail MLIR op creation at runtime even though the import succeeded (observed:
|
||||
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
|
||||
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
|
||||
class _FA4Policy:
|
||||
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
|
||||
|
||||
FA4 is skipped when:
|
||||
* a previous call failed at runtime -- CuTeDSL JIT compilation is
|
||||
shape-dependent, so the first failure disables FA4 for the rest of
|
||||
the process instead of retrying a broken JIT on every call; or
|
||||
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
|
||||
dies on its arch check) and is unvalidated for training in this repo
|
||||
(its lse is not even allocated through our inference-shaped custom
|
||||
op), so training keeps the pre-FA4 behavior: FA2 on every device.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.broken = False
|
||||
|
||||
def use_fa4(self, *tensors: torch.Tensor) -> bool:
|
||||
if self.broken:
|
||||
return False
|
||||
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
|
||||
|
||||
def mark_broken(self, error: Exception) -> None:
|
||||
if not self.broken:
|
||||
self.broken = True
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
|
||||
"to FA2 for the rest of this process.", error)
|
||||
|
||||
|
||||
_FA4 = _FA4Policy()
|
||||
|
||||
|
||||
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
|
||||
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
|
||||
|
||||
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
|
||||
the first FA4 runtime failure) the call is served by ``fa2_func``.
|
||||
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
|
||||
failure, so it propagates without disabling FA4.
|
||||
"""
|
||||
|
||||
def decorator(fa4_func: Callable) -> Callable:
|
||||
|
||||
@functools.wraps(fa4_func)
|
||||
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
if _FA4.use_fa4(q, k, v):
|
||||
try:
|
||||
return fa4_func(q, k, v, *args, **kwargs)
|
||||
except NotImplementedError:
|
||||
raise
|
||||
except Exception as e: # CuTeDSL compile errors surface as ValueError
|
||||
_FA4.mark_broken(e)
|
||||
return fa2_func(q, k, v, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_func)
|
||||
def flash_attn_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -317,6 +287,16 @@ def flash_attn_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
|
||||
return out
|
||||
@@ -392,7 +372,6 @@ def flash_attn_fp4_func(
|
||||
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_varlen_func)
|
||||
def flash_attn_varlen_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -407,6 +386,20 @@ def flash_attn_varlen_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_varlen_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
|
||||
q,
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""torch.compile-traceable wrapper for the FA2/FA3/FA4 default attention path.
|
||||
|
||||
The FA4/cute path (`fa_version == "4"`) is already a registered
|
||||
`torch.library.custom_op` in `fastvideo.attention.utils.flash_attn_cute`, so
|
||||
dynamo treats it as a graph node. The external FA2/FA3 ``flash_attn_func`` is
|
||||
NOT — dynamo breaks the graph at the call site (observed: wanvideo.py
|
||||
self-attn, once per layer every step), which fragments the compiled region
|
||||
and blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom op
|
||||
(mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
(correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
boundary changes, so numerics are unchanged (SSIM-gated).
|
||||
|
||||
Autograd: FA2 has full ``register_autograd`` parity — the custom op's
|
||||
backward calls flash_attn's ``_flash_attn_backward`` directly, so training
|
||||
backprops *through* the op (no graph break on the training path either).
|
||||
FA3 currently keeps the no-backward + carve-out pattern from PR #1373
|
||||
because FA3's private backward signature wants validation on a real Hopper
|
||||
box (gated on Kuan-Hao's Modal FA3 setup PR). Once that lands the FA3 path
|
||||
can mirror FA2.
|
||||
|
||||
Lives in `attention/utils/` (sibling of `flash_attn_cute.py` and
|
||||
`flash_attn_no_pad.py`) so it can be imported by any backend that wants the
|
||||
traceable FA default call without pulling in backend dispatch logic. The
|
||||
backend (`attention/backends/flash_attn.py`) just imports
|
||||
`flash_attn_func_compilable` and `fa_version` from here.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Pick the same backend the rest of FastVideo picked for `flash_attn_func`
|
||||
# (FA4/cute → FA3 → FA2). Mirror the precedence used in
|
||||
# `attention/utils/flash_attn_no_pad.py` so the two probes always agree.
|
||||
#
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1: its CuTeDSL
|
||||
# kernels JIT-compile per shape family and can fail at runtime on some
|
||||
# arch/shape combinations, so it is never auto-selected just because it is
|
||||
# installed.
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer has a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if fa_version == "2":
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal call
|
||||
# shape used by FlashAttentionImpl.forward's default branch (see
|
||||
# `flash_attn_func_compilable(...)` call site in
|
||||
# `attention/backends/flash_attn.py`). The masked/no-pad and varlen /
|
||||
# cross-attn paths use different entry points
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_*`) which live in
|
||||
# `attention/utils/flash_attn_no_pad.py`. The wrapper's signature is the
|
||||
# contract: any extra kwarg (dropout_p, window_size, alibi_slopes,
|
||||
# deterministic, return_attn_probs, ...) raises TypeError at the call
|
||||
# site, so silent loss of kwargs is not a failure mode.
|
||||
from flash_attn.flash_attn_interface import _flash_attn_backward as _fa2_backward
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# `return_attn_probs=True` asks FA2 to also return softmax_lse +
|
||||
# S_dmask. We need softmax_lse to feed the backward; S_dmask is the
|
||||
# dropout mask (always None here since dropout_p is fixed at 0).
|
||||
out, softmax_lse, _ = _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal, return_attn_probs=True)
|
||||
return out, softmax_lse
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del softmax_scale, causal
|
||||
# FA2 default path: out = [batch, seqlen_q, nheads, head_dim_v],
|
||||
# softmax_lse = [batch, nheads, seqlen_q], fp32 regardless of q dtype.
|
||||
b, sq, hq = q.shape[0], q.shape[1], q.shape[2]
|
||||
out = q.new_empty(b, sq, hq, v.shape[-1])
|
||||
lse = q.new_empty(b, hq, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_default_setup_context(ctx, inputs, output):
|
||||
q, k, v, softmax_scale, causal = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(q, k, v, out, lse)
|
||||
# `lse` is an auxiliary output we save to feed FA2's backward; nobody
|
||||
# should differentiate through it. Mark it non-differentiable so
|
||||
# autograd errors loudly if a caller wires it into a loss, rather
|
||||
# than silently producing zero/None grads through the `del grad_lse`
|
||||
# in our backward.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's *forward* substitutes `1 / sqrt(head_dim)` for `softmax_scale=None`
|
||||
# internally; FA2's *backward* (`_flash_attn_backward`) demands a concrete
|
||||
# float in its C++ schema and rejects None at the binding boundary. Resolve
|
||||
# the default here so the value saved on ctx (and passed to backward) is
|
||||
# always a real float — matches what FA2's own autograd.Function does.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = q.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
|
||||
def _flash_attn_default_backward(ctx, grad_out, grad_lse):
|
||||
# We only differentiate `out`; softmax_lse is saved-for-backward, not
|
||||
# a real differentiable output. (Mirrors the FP4 cute template.)
|
||||
del grad_lse
|
||||
q, k, v, out, lse = ctx.saved_tensors
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
# FA2's `_flash_attn_backward` writes into dq/dk/dv in place. The
|
||||
# extra kwargs (window_size_*, softcap, alibi_slopes, deterministic,
|
||||
# rng_state) are pinned to the same defaults the forward wrapper
|
||||
# uses — flash-attn==2.8.1 (the version FastVideo pins) requires
|
||||
# all of them explicitly. `rng_state=None` is correct for our
|
||||
# `dropout_p=0` configuration.
|
||||
_fa2_backward(
|
||||
grad_out,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
out,
|
||||
lse,
|
||||
dq,
|
||||
dk,
|
||||
dv,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=False,
|
||||
rng_state=None,
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
_flash_attn_default_backward,
|
||||
setup_context=_flash_attn_default_setup_context,
|
||||
)
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Backward is registered: autograd flows through the op (training
|
||||
# path is also traceable; no carve-out needed). Public API matches
|
||||
# `flash_attn_func` — returns just `out`; we drop the saved-for-
|
||||
# backward `lse` here so callers see the original single-tensor
|
||||
# contract.
|
||||
out, _ = torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
return out
|
||||
elif fa_version == "3":
|
||||
# FA3 path: same forward+fake custom op as the original PR #1373, with
|
||||
# the autograd carve-out kept. The full backward (mirroring the FA2 leg
|
||||
# above) wants a Hopper box for grad-check validation, which we don't
|
||||
# have until Kuan-Hao's Modal FA3 setup PR lands. Until then this keeps
|
||||
# inference traceable + training correct (via the original
|
||||
# autograd.Function path + a pre-PR-style graph break on training).
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
@@ -21,27 +21,46 @@ from einops import rearrange
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from fastvideo import envs
|
||||
|
||||
def _resolve_flash_attn_varlen_func() -> Any:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
|
||||
return flash_attn_varlen_func_cute
|
||||
except ImportError:
|
||||
def _resolve_flash_attn_varlen_func() -> tuple[Any, str]:
|
||||
if envs.FASTVIDEO_FA4:
|
||||
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
|
||||
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
|
||||
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
|
||||
# not ImportError: importers like bsa_attn.py treat ImportError as
|
||||
# "flash-attn not installed" and silently degrade to reference kernels.
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
return flash_attn_varlen_func_cute, "4"
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_flash
|
||||
return flash_attn_varlen_func_interface, "3"
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_flash, "2"
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
flash_attn_varlen_func_impl, _FA_VARLEN_VERSION = _resolve_flash_attn_varlen_func()
|
||||
|
||||
# FA2-only: the private varlen backward we register against the custom ops
|
||||
# below. FA3 / FA4 have different private signatures and validation paths
|
||||
# (Hopper / Blackwell boxes) — those legs keep the autograd carve-out
|
||||
# pattern from PR #1373 until their setup PRs land.
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
from flash_attn.flash_attn_interface import (
|
||||
_flash_attn_varlen_backward as _fa2_varlen_backward, )
|
||||
|
||||
|
||||
def flash_attn_no_pad(
|
||||
@@ -180,3 +199,473 @@ def flash_attn_varlen_qk_no_pad(
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# torch.compile traceability + register_autograd parity for the masked /
|
||||
# varlen attention paths.
|
||||
#
|
||||
# Wraps the two entry points `FlashAttentionImpl.forward` calls
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_qk_no_pad`) as
|
||||
# `torch.library.custom_op`s so dynamo sees one traceable node — the
|
||||
# internal unpad / pad bookkeeping (data-dependent `nnz` shapes) runs
|
||||
# eager inside the op, and the op's outputs are the statically-shaped
|
||||
# padded tensors. This mirrors the FA2 default-path wrapper in
|
||||
# `fastvideo/attention/backends/flash_attn.py`.
|
||||
#
|
||||
# Autograd: on FA2 we register a real backward (`register_autograd`)
|
||||
# that calls FA2's `_flash_attn_varlen_backward` on the unpadded form
|
||||
# — re-unpadding the saved padded tensors using the saved mask. The
|
||||
# `softmax_lse` from the varlen forward is naturally unpadded
|
||||
# (`[nheads, total_q]`); we pad it to `[batch, nheads, seqlen]` on
|
||||
# the way out (statically shaped) and re-unpad in backward. So
|
||||
# training backprops *through* the op (no graph break on the training
|
||||
# path either).
|
||||
#
|
||||
# FA3 / FA4 keep the autograd carve-out pattern from PR #1373: the
|
||||
# custom op has forward + fake only, and `*_compilable` falls back to
|
||||
# the original autograd.Function for grad-enabled calls. Those legs
|
||||
# are gated on Hopper-class / Blackwell-class boxes for backward
|
||||
# validation and ship as separate follow-ups.
|
||||
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
# ---------- masked self-attention: flash_attn_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, s, _three, h, d = qkv.shape
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_qkvpacked_func(x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz, h, d] -> [b, s, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), indices, b, s),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: FA2 varlen returns [nheads, total_q]. Transpose to [total_q,
|
||||
# nheads], pad to [b, s, nheads], permute to [b, nheads, s] — statically
|
||||
# shaped so register_fake matches.
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), indices, b, s).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
out = qkv.new_empty(b, s, h, d)
|
||||
lse = qkv.new_empty(b, h, s, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_no_pad_setup_context(ctx, inputs, output):
|
||||
qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(qkv, out, lse, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's varlen backward requires a concrete float for softmax_scale.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = qkv.shape[-1]**-0.5 # head_dim from qkv's last dim
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
# lse is saved-for-backward, not differentiated.
|
||||
del grad_lse
|
||||
qkv, out_padded, lse_padded, key_padding_mask = ctx.saved_tensors
|
||||
b, s, _three, h, d = qkv.shape
|
||||
|
||||
# One `unpad_input` call (on qkv) gives us indices + cu_seqlens + max_s;
|
||||
# reuse those for out / dout / lse below via direct indexing instead
|
||||
# of redundant `unpad_input` calls (each of which would re-run
|
||||
# `nonzero` + `cumsum` + a `.max().item()` GPU→CPU sync).
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
q_unpad, k_unpad, v_unpad = (t.contiguous() for t in x_unpad.unbind(dim=1))
|
||||
|
||||
# Direct-index variants reuse `indices` (computed above).
|
||||
out_unpad = out_padded.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
# lse_padded [b, h, s] -> [b, s, h] -> [nnz, h] -> [h, nnz].
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens,
|
||||
cu_seqlens_k=cu_seqlens,
|
||||
max_seqlen_q=max_s,
|
||||
max_seqlen_k=max_s,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# Re-pad each grad and stack into dqkv.
|
||||
def _repad(dt_unpad: torch.Tensor) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, b, s)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dqkv = torch.stack([_repad(dq_unpad), _repad(dk_unpad), _repad(dv_unpad)], dim=2)
|
||||
# 6 inputs total: qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic.
|
||||
return dqkv, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
_flash_attn_no_pad_backward,
|
||||
setup_context=_flash_attn_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- cross-attention: flash_attn_varlen_qk_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, sq, h, d = query.shape
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
v_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
v_unpad = rearrange(v_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_func_impl(q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz_q, h, d] -> [b, sq, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), q_indices, b, sq),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: [h, nnz_q] -> [b, h, sq]
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), q_indices, b, sq).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# out_padded ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
out = query.new_empty(b, sq, h, value.shape[-1])
|
||||
lse = query.new_empty(b, h, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_setup_context(ctx, inputs, output):
|
||||
(query, key, value, query_padding_mask, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic) = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(query, key, value, out, lse, query_padding_mask, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
if softmax_scale is None:
|
||||
softmax_scale = query.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
del grad_lse
|
||||
(query, key, value, out_padded, lse_padded, query_padding_mask, key_padding_mask) = ctx.saved_tensors
|
||||
b, sq, h, d = query.shape
|
||||
sk = key.shape[1]
|
||||
|
||||
# One `unpad_input` call per distinct mask; reuse the returned
|
||||
# indices via direct indexing for everything else that shares
|
||||
# the same mask (v with k_mask; out/dout/lse with q_mask; the
|
||||
# final repad of dk/dv also reuses k_indices). Avoids ~4
|
||||
# redundant `unpad_input` calls + their GPU→CPU `.max().item()`
|
||||
# syncs.
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, k_indices, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
v_unpad = value.flatten(0, 1)[k_indices].view(-1, h, d).contiguous()
|
||||
|
||||
# out / dout / lse follow q's shape, so index with q_indices.
|
||||
out_unpad = out_padded.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[q_indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# k_indices is already available from the unpad_input above —
|
||||
# no need to recompute it for the dk/dv repad.
|
||||
def _repad(dt_unpad: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, batch, seqlen)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dq_padded = _repad(dq_unpad, q_indices, b, sq)
|
||||
dk_padded = _repad(dk_unpad, k_indices, b, sk)
|
||||
dv_padded = _repad(dv_unpad, k_indices, b, sk)
|
||||
# 9 inputs total: query, key, value, q_mask, k_mask, causal, dropout_p,
|
||||
# softmax_scale, deterministic.
|
||||
return dq_padded, dk_padded, dv_padded, None, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
_flash_attn_varlen_qk_no_pad_backward,
|
||||
setup_context=_flash_attn_varlen_qk_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- public dispatchers (FA2: autograd flows through the op) -----
|
||||
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_no_pad`` (registered op,
|
||||
full register_autograd on FA2 — both inference and training go through
|
||||
the op, no graph break on either)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
return out
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_varlen_qk_no_pad`` (registered
|
||||
op, full register_autograd on FA2)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
return out
|
||||
|
||||
else:
|
||||
# ---------- FA3 / FA4: carve-out (forward+fake only, no real backward) ---
|
||||
# Same pattern as the parked varlen-extension and the FA3 default leg in
|
||||
# `fastvideo/attention/backends/flash_attn.py`. Real backward for these
|
||||
# versions is a follow-up gated on Hopper / Blackwell box validation.
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_no_pad( # type: ignore[no-untyped-call]
|
||||
qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
return qkv.new_empty(b, s, h, d)
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_varlen_qk_no_pad( # type: ignore[no-untyped-call]
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# output ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
return query.new_empty(b, sq, h, value.shape[-1])
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and qkv.requires_grad:
|
||||
return flash_attn_no_pad(qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad):
|
||||
return flash_attn_varlen_qk_no_pad(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
@@ -13,7 +16,8 @@ from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldArchConfig(WanVideoArchConfig):
|
||||
"""DreamX-World DiT config with camera PRoPE control fields."""
|
||||
|
||||
add_control_adapter: bool = True
|
||||
cam_method: str | None = "prope"
|
||||
attn_compress: int = 1
|
||||
cam_self_attn_layers: tuple[int, ...] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DreamXWorldArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldARArchConfig(DreamXWorldArchConfig):
|
||||
"""DreamX-World-5B autoregressive causal DiT config."""
|
||||
|
||||
model_type: str = "ti2v"
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
text_len: int = 512
|
||||
text_dim: int = 4096
|
||||
freq_dim: int = 256
|
||||
attn_compress: int = 4
|
||||
cam_self_attn_layers: tuple[int, ...] | None = tuple(range(30))
|
||||
local_attn_size: int = 12
|
||||
sink_size: int = 3
|
||||
num_frames_per_block: int = 3
|
||||
rope_cache_policy: str = "block_relativistic"
|
||||
# The official AR checkpoint (AMAP-ML/DreamX-World ``model.safetensors``)
|
||||
# already uses FastVideo's native key names and the converter copies the
|
||||
# tensors verbatim, so every rule is an identity. The rules enumerate the
|
||||
# full state-dict surface of ``DreamXWorldARTransformer3DModel`` (norm1 /
|
||||
# norm2 / head.norm are affine-free and have no parameters).
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^patch_embedding\.(.*)$": r"patch_embedding.\1",
|
||||
r"^text_embedding\.([02])\.(.*)$": r"text_embedding.\1.\2",
|
||||
r"^time_embedding\.([02])\.(.*)$": r"time_embedding.\1.\2",
|
||||
r"^time_projection\.1\.(.*)$": r"time_projection.1.\1",
|
||||
r"^blocks\.(\d+)\.self_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.self_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_(q|k)\.weight$": r"blocks.\1.self_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.cross_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.cross_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_(q|k)\.weight$": r"blocks.\1.cross_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.cam_self_attn\.(q_proj|k_proj|v_proj|out_proj)\.(.*)$": r"blocks.\1.cam_self_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.cam_self_attn\.norm_(q|k)\.weight$": r"blocks.\1.cam_self_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm3.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.([02])\.(.*)$": r"blocks.\1.ffn.\2.\3",
|
||||
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.modulation",
|
||||
r"^head\.head\.(.*)$": r"head.head.\1",
|
||||
r"^head\.modulation$": r"head.modulation",
|
||||
})
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldARConfig(DreamXWorldConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DreamXWorldARArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxTransformer2DArchConfig(DiTArchConfig):
|
||||
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = None
|
||||
num_layers: int = 19
|
||||
num_single_layers: int = 38
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 24
|
||||
joint_attention_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
guidance_embeds: bool = True
|
||||
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=FluxTransformer2DArchConfig)
|
||||
prefix: str = "flux"
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTArchConfig(DiTArchConfig):
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
hidden_size: int = 4096
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_layers: int = 30
|
||||
|
||||
text_embed_dim: int = 1472
|
||||
time_embed_dim: int = 512
|
||||
condition_dim: int = 256
|
||||
|
||||
prior_vq_quantizer_codebook_size: int = 16384
|
||||
|
||||
patch_size: int = 2
|
||||
|
||||
max_height: int = 2048
|
||||
max_width: int = 2048
|
||||
|
||||
qk_norm: str = "layer_norm"
|
||||
eps: float = 1e-5
|
||||
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
|
||||
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
|
||||
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
|
||||
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.num_channels_latents = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
|
||||
prefix: str = "GlmImage"
|
||||
@@ -2,14 +2,24 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_kandinsky5_transformer_block(n: str, m) -> bool:
|
||||
return ("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
lambda n, m:
|
||||
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_kandinsky5_transformer_block])
|
||||
|
||||
# NABLA block-sparse attention for attention_type="nabla" checkpoints, plus
|
||||
# the dense backends every DiT supports.
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.NABLA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
# Native FastVideo implementation uses the same parameter names as diffusers
|
||||
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -19,6 +20,11 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
# AnyFlow dual-timestep checkpoints expose delta_embedder weights with the
|
||||
# same internal layout as time_embedder. The regex is harmless on plain
|
||||
# Wan checkpoints (no delta_embedder keys to match).
|
||||
r"^condition_embedder\.delta_embedder\.linear_1\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.delta_embedder\.linear_2\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
@@ -86,6 +92,14 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
|
||||
rope_cache_policy: str = "absolute"
|
||||
|
||||
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
|
||||
# the legacy single-timestep forward (no delta_embedder allocated, no
|
||||
# extra computation on the embedder forward path).
|
||||
r_embedder: bool = False
|
||||
r_embedder_fusion: Literal["additive", "gated"] = "additive"
|
||||
r_embedder_gate_value: float = 0.25
|
||||
r_embedder_deltatime_type: Literal["r", "t-r"] = "r"
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
|
||||
@@ -43,11 +43,16 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
require_processor: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
# update_model_arch re-runs __post_init__ after pipeline configs may
|
||||
# have customized tokenizer_kwargs (e.g. kandinsky5/gen3c/longcat set
|
||||
# "padding"); rebuilding the dict here would silently wipe those
|
||||
# customizations, so only fill in defaults for keys not already set.
|
||||
defaults = {
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
self.tokenizer_kwargs = defaults | self.tokenizer_kwargs
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
|
||||
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
@@ -21,4 +22,5 @@ __all__ = [
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
"Flux2VAEConfig",
|
||||
"GlmImageVAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import (AutoencoderKLArchConfig, AutoencoderKLVAEConfig)
|
||||
|
||||
_GLM_IMAGE_LATENTS_MEAN: tuple[float, ...] = (
|
||||
-0.2080078125,
|
||||
1.875,
|
||||
-0.470703125,
|
||||
-1.265625,
|
||||
-1.421875,
|
||||
0.77734375,
|
||||
-0.3671875,
|
||||
-0.9453125,
|
||||
0.318359375,
|
||||
0.7734375,
|
||||
-0.1884765625,
|
||||
-0.022216796875,
|
||||
-0.220703125,
|
||||
-1.59375,
|
||||
-0.81640625,
|
||||
-0.255859375,
|
||||
)
|
||||
_GLM_IMAGE_LATENTS_STD: tuple[float, ...] = (
|
||||
3.0625,
|
||||
2.203125,
|
||||
2.265625,
|
||||
4.84375,
|
||||
2.5,
|
||||
3.9375,
|
||||
2.203125,
|
||||
3.03125,
|
||||
2.1875,
|
||||
2.046875,
|
||||
2.71875,
|
||||
2.390625,
|
||||
2.390625,
|
||||
2.453125,
|
||||
2.25,
|
||||
2.15625,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEArchConfig(AutoencoderKLArchConfig):
|
||||
act_fn: str = "silu"
|
||||
block_out_channels: tuple[int, ...] = (128, 512, 1024, 1024)
|
||||
down_block_types: tuple[str, ...] = (
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
)
|
||||
up_block_types: tuple[str, ...] = (
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
)
|
||||
force_upcast: bool = True
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 16
|
||||
latents_mean: tuple[float, ...] = _GLM_IMAGE_LATENTS_MEAN
|
||||
latents_std: tuple[float, ...] = _GLM_IMAGE_LATENTS_STD
|
||||
layers_per_block: int = 3
|
||||
mid_block_add_attention: bool = False
|
||||
norm_num_groups: int = 32
|
||||
sample_size: int = 1024
|
||||
scaling_factor: float = 0.18215
|
||||
shift_factor: float | None = None
|
||||
use_quant_conv: bool = False
|
||||
use_post_quant_conv: bool = False
|
||||
|
||||
temporal_compression_ratio: int = 1
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEConfig(AutoencoderKLVAEConfig):
|
||||
arch_config: GlmImageVAEArchConfig = field(default_factory=GlmImageVAEArchConfig)
|
||||
|
||||
use_tiling: bool = True
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
tile_sample_min_height: int = 512
|
||||
tile_sample_min_width: int = 512
|
||||
tile_sample_stride_height: int = 384
|
||||
tile_sample_stride_width: int = 384
|
||||
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = True
|
||||
@@ -1,10 +1,12 @@
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
@@ -16,5 +18,6 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
|
||||
"Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B-Cam FastVideo model configuration helpers."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import (DreamXWorldARArchConfig, DreamXWorldARConfig,
|
||||
DreamXWorldArchConfig, DreamXWorldConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.wan import LucyEditDevConfig, t5_postprocess_text
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_dit_config() -> DreamXWorldConfig:
|
||||
"""Return the DreamX-World DiT config matching DreamX-World-5B-Cam."""
|
||||
return DreamXWorldConfig(arch_config=DreamXWorldArchConfig(
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
in_channels=48,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
cross_attn_norm=True,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=None,
|
||||
))
|
||||
|
||||
|
||||
def make_dreamx_world_5b_ar_dit_config() -> DreamXWorldARConfig:
|
||||
"""Return the DreamX-World-5B autoregressive causal DiT config."""
|
||||
return DreamXWorldARConfig(arch_config=DreamXWorldARArchConfig(
|
||||
model_type="ti2v",
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
in_channels=48,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
cross_attn_norm=True,
|
||||
qk_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=4,
|
||||
cam_self_attn_layers=tuple(range(30)),
|
||||
local_attn_size=12,
|
||||
sink_size=3,
|
||||
num_frames_per_block=3,
|
||||
))
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_vae_config() -> WanVAEConfig:
|
||||
"""Return the Wan2.2 48-channel VAE config used by DreamX-World-5B-Cam."""
|
||||
return LucyEditDevConfig().vae_config
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_text_encoder_config() -> T5Config:
|
||||
"""Return the UMT5-XXL text encoder config used by DreamX-World-5B-Cam."""
|
||||
return T5Config(
|
||||
arch_config=T5ArchConfig(
|
||||
vocab_size=256384,
|
||||
d_model=4096,
|
||||
d_kv=64,
|
||||
d_ff=10240,
|
||||
num_layers=24,
|
||||
num_decoder_layers=None,
|
||||
num_heads=64,
|
||||
relative_attention_num_buckets=32,
|
||||
dropout_rate=0.0,
|
||||
text_len=512,
|
||||
feed_forward_proj="gelu",
|
||||
is_encoder_decoder=False,
|
||||
),
|
||||
prefix="umt5",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorld5BCamPipelineConfig(PipelineConfig):
|
||||
"""Pipeline config for the first-scope DreamX-World-5B-Cam mode."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_cam_dit_config)
|
||||
vae_config: VAEConfig = field(default_factory=make_dreamx_world_5b_cam_vae_config)
|
||||
text_encoder_configs: tuple[EncoderConfig,
|
||||
...] = field(default_factory=lambda: (make_dreamx_world_5b_cam_text_encoder_config(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (t5_postprocess_text, ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
flow_shift: float | None = 3.0
|
||||
ti2v_task: bool = True
|
||||
expand_timesteps: bool = True
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
vae_precision: str = "fp32"
|
||||
vae_decode_precision: str | None = "bf16"
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorld5BARPipelineConfig(DreamXWorld5BCamPipelineConfig):
|
||||
"""Pipeline config for DreamX-World-5B autoregressive forcing."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_ar_dit_config)
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
is_causal: bool = True
|
||||
dmd_denoising_steps: tuple[int, ...] = (1000, 750, 500, 250)
|
||||
warp_denoising_step: bool = True
|
||||
context_noise: float = 0.1
|
||||
num_frames_per_block: int = 3
|
||||
color_correction_strength: float = 1.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.dit_config.expand_timesteps = True
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import EncoderConfig
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.encoders import (
|
||||
BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
T5LargeConfig,
|
||||
)
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def _flux_clip_pooled_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""CLIP branch for FLUX: Diffusers uses pooled prompt embeddings only."""
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError(
|
||||
"FLUX CLIP conditioning requires pooler_output. Ensure the CLIP text encoder returns pooled features.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
def _flux_t5_sequence_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.last_hidden_state is None:
|
||||
raise RuntimeError("FLUX T5 conditioning requires last_hidden_state.")
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxPipelineConfig(PipelineConfig):
|
||||
"""Pipeline layout for Diffusers FLUX.1-dev (CLIP + T5 + packed DiT + FlowMatch)."""
|
||||
|
||||
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
|
||||
transformer_arch: str = "FluxTransformer2DModel"
|
||||
vae_arch: str = "AutoencoderKL"
|
||||
text_encoder_archs: tuple[str, ...] = ("CLIPTextModel", "T5EncoderModel")
|
||||
tokenizer_archs: tuple[str, ...] = ("CLIPTokenizer", "T5TokenizerFast")
|
||||
|
||||
dit_config: FluxDiTConfig = field(default_factory=FluxDiTConfig)
|
||||
vae_config: AutoencoderKLVAEConfig = field(default_factory=AutoencoderKLVAEConfig)
|
||||
|
||||
embedded_cfg_scale: float = 3.5
|
||||
flow_shift: float | None = None
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (CLIPTextConfig(), T5LargeConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str],
|
||||
...] = field(default_factory=lambda: (preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(
|
||||
default_factory=lambda: (_flux_clip_pooled_postprocess, _flux_t5_sequence_postprocess))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", "bf16"))
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
te_cfgs = list(self.text_encoder_configs)
|
||||
if len(te_cfgs) >= 1:
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("max_length", 77)
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
if len(te_cfgs) >= 2:
|
||||
cap = 512
|
||||
te_cfgs[1].tokenizer_kwargs["max_length"] = min(int(te_cfgs[1].tokenizer_kwargs.get("max_length", cap)),
|
||||
cap)
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def glm_image_t5_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
assert torch.isnan(hidden_state).sum() == 0, "T5 hidden states contain NaN"
|
||||
|
||||
max_len = 512
|
||||
prompt_embeds = [u[:min(v, max_len)] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_len - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0)
|
||||
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageConfig(PipelineConfig):
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=GlmImageDiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig)
|
||||
vae_precision: str = "fp32"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (glm_image_t5_postprocess, ))
|
||||
|
||||
flow_shift: float | None = 1.0
|
||||
embedded_cfg_scale: float = 7.5
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, CLIPTextConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1Config
|
||||
from fastvideo.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
# Byte-exact copy of the upstream Kandinsky5/diffusers template, including the
|
||||
# "promt"/"scren" typos: the checkpoints were trained with this exact system # codespell:ignore promt,scren
|
||||
# prompt, and ENCODE_START_IDX below is the tokenized length of everything
|
||||
# before the user prompt. Fixing the typos shifts user content to index 127
|
||||
# and mis-conditions every generation.
|
||||
KANDINSKY5_PROMPT_TEMPLATE = "\n".join([
|
||||
"<|im_start|>system\nYou are a promt engineer. Describe the video in detail.", # codespell:ignore promt
|
||||
"Describe how the camera moves or shakes, describe the zoom and view angle, whether it follows the objects.",
|
||||
"Describe the location of the video, main characters or objects and their action.",
|
||||
"Describe the dynamism of the video and presented actions.",
|
||||
"Name the visual style of the video: whether it is a professional footage, user generated content, some kind of animation, video game or scren content.", # codespell:ignore scren
|
||||
"Describe the visual effects, postprocessing and transitions if they are presented in the video.",
|
||||
"Pay attention to the order of key actions shown in the scene.<|im_end|>",
|
||||
"<|im_start|>user\n{}<|im_end|>",
|
||||
])
|
||||
KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX = 129
|
||||
|
||||
|
||||
def kandinsky5_qwen_preprocess_text(prompt: str) -> str:
|
||||
if not prompt.strip():
|
||||
prompt = "."
|
||||
return KANDINSKY5_PROMPT_TEMPLATE.format(prompt)
|
||||
|
||||
|
||||
def kandinsky5_qwen_postprocess_text(outputs: BaseEncoderOutput,
|
||||
mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if outputs.hidden_states is None:
|
||||
raise RuntimeError("Kandinsky5 Qwen prompt embeddings require hidden_states.")
|
||||
hidden_states = outputs.hidden_states[-1]
|
||||
prompt_embeds = hidden_states[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
mask = mask[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
if prompt_embeds.shape[1] == 0:
|
||||
prompt_embeds = hidden_states[:, -1:]
|
||||
mask = torch.ones((mask.shape[0], 1), dtype=mask.dtype, device=mask.device)
|
||||
return prompt_embeds, mask
|
||||
|
||||
|
||||
def kandinsky5_clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError("Kandinsky5 CLIP pooled output is required.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5T2VConfig(PipelineConfig):
|
||||
"""Kandinsky-5.0 Lite text-to-video pipeline configuration."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Kandinsky5VideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Reason1Config(), CLIPTextConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_postprocess_text, kandinsky5_clip_postprocess_text))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
|
||||
text_encoder_max_lengths: tuple[int, ...] = field(
|
||||
default_factory=lambda: (KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX + 512, 77))
|
||||
|
||||
flow_shift: float | None = 5.0
|
||||
vae_tiling: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if len(self.text_encoder_configs) != 2:
|
||||
raise ValueError(f"Kandinsky5 pipeline requires exactly 2 text encoders (qwen and clip), "
|
||||
f"but got {len(self.text_encoder_configs)} encoder(s).")
|
||||
if len(self.text_encoder_precisions) != 2:
|
||||
raise ValueError("Kandinsky5 pipeline requires exactly 2 text encoder precisions, "
|
||||
f"but got {len(self.text_encoder_precisions)}.")
|
||||
if len(self.text_encoder_max_lengths) != 2:
|
||||
raise ValueError("Kandinsky5 pipeline requires exactly 2 text encoder max lengths, "
|
||||
f"but got {len(self.text_encoder_max_lengths)}.")
|
||||
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
qwen_cfg = self.text_encoder_configs[0]
|
||||
qwen_cfg.arch_config.output_hidden_states = True
|
||||
qwen_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": True,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
clip_cfg = self.text_encoder_configs[1]
|
||||
clip_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": "max_length",
|
||||
"max_length": 77,
|
||||
"truncation": True,
|
||||
"add_special_tokens": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5I2VConfig(Kandinsky5T2VConfig):
|
||||
"""Kandinsky-5.0 image-to-video pipeline configuration."""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# I2V needs the VAE encoder to encode the conditioning image.
|
||||
self.vae_config.load_encoder = True
|
||||
@@ -2,6 +2,7 @@
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
@@ -94,8 +95,56 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
|
||||
return len(self.sp_group_local_indices) // self.batch_size
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str):
|
||||
dataset_root = os.path.realpath(os.path.expanduser(path))
|
||||
def _parse_data_path_specs(path: str | Sequence[str] | dict[str, int]) -> list[tuple[str, int]]:
|
||||
"""Parse one or more dataset roots with old-framework repeat counts."""
|
||||
if isinstance(path, dict):
|
||||
return [
|
||||
(str(root), int(repeat))
|
||||
for root, repeat in path.items()
|
||||
if int(repeat) > 0
|
||||
]
|
||||
if isinstance(path, Sequence) and not isinstance(path, str):
|
||||
return [(str(root), 1) for root in path]
|
||||
|
||||
specs: list[tuple[str, int]] = []
|
||||
for part in str(path).split(","):
|
||||
part = part.strip()
|
||||
if not part:
|
||||
continue
|
||||
if ":" in part:
|
||||
dir_path, count_str = part.rsplit(":", 1)
|
||||
count = int(count_str)
|
||||
else:
|
||||
dir_path, count = part, 1
|
||||
if count > 0:
|
||||
specs.append((dir_path.strip(), count))
|
||||
return specs
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]):
|
||||
specs = _parse_data_path_specs(path)
|
||||
if len(specs) != 1 or specs[0][1] != 1:
|
||||
all_file_names: list[str] = []
|
||||
all_lengths: list[int] = []
|
||||
for root, repeat in specs:
|
||||
file_names, lengths = get_parquet_files_and_length(str(root))
|
||||
for _ in range(repeat):
|
||||
all_file_names.extend(file_names)
|
||||
all_lengths.extend(lengths)
|
||||
if not all_file_names:
|
||||
raise FileNotFoundError(
|
||||
"No parquet files found under dataset paths: "
|
||||
f"{path}. "
|
||||
"Please verify these paths point to preprocessed parquet data."
|
||||
)
|
||||
file_lengths = sorted(
|
||||
zip(all_file_names, all_lengths, strict=True),
|
||||
key=lambda x: x[0],
|
||||
)
|
||||
file_names_sorted, lengths_sorted = zip(*file_lengths, strict=True)
|
||||
return file_names_sorted, lengths_sorted
|
||||
|
||||
dataset_root = os.path.realpath(os.path.expanduser(specs[0][0]))
|
||||
# Check if cached info exists
|
||||
cache_dir = os.path.join(dataset_root, "map_style_cache")
|
||||
cache_file = os.path.join(cache_dir, "file_info.pkl")
|
||||
@@ -268,7 +317,7 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
path: str | Sequence[str] | dict[str, int],
|
||||
batch_size: int,
|
||||
parquet_schema: pa.Schema,
|
||||
cfg_rate: float = 0.0,
|
||||
@@ -358,7 +407,6 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
def __len__(self):
|
||||
return sum(self.lengths)
|
||||
|
||||
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
# 3. Loader helper – everything else stays just like your original trainer
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -23,6 +23,7 @@ If you only need to use the distributed environment without model parallelism,
|
||||
you can skip the model parallel initialization and destruction steps.
|
||||
"""
|
||||
import contextlib
|
||||
import gc
|
||||
import os
|
||||
import pickle
|
||||
import weakref
|
||||
@@ -1006,6 +1007,13 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
if shutdown_ray:
|
||||
import ray # Lazy import Ray
|
||||
ray.shutdown()
|
||||
# Actually free GPU memory, as the name promises. FSDP-wrapped modules sit
|
||||
# in reference cycles (module <-> FSDP state <-> hooks), so dropping the
|
||||
# last user reference does not free their parameters until a gc pass runs.
|
||||
# Without this, back-to-back model-load tests in one pytest session OOM.
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def get_same_node_ranks(pg: ProcessGroup | StatelessProcessGroup, source_rank: int = 0) -> list[int]:
|
||||
|
||||
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_FA4: bool = False
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
@@ -207,9 +208,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
|
||||
# FASTVIDEO_FA4=1 as well (see below).
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
|
||||
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
|
||||
# is installed. Below sm90, grad-enabled and GQA calls are routed to FA2
|
||||
# (FA4's backward asserts sm90+ and its pack_gqa fails to JIT there).
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
@@ -122,6 +122,44 @@ to run a subset of the Evaluator's registered metrics on this batch
|
||||
(useful for scoring different corpora with different metric subsets in
|
||||
successive calls without burning model loads).
|
||||
|
||||
### Training-time validation metrics
|
||||
|
||||
The modular trainer can score validation videos during training through
|
||||
`callbacks.validation.metrics`. Each distributed rank that writes local
|
||||
validation videos also runs its local metrics on `cuda:<local_rank>`;
|
||||
rank 0 only merges scalar summaries and logs artifacts.
|
||||
|
||||
```yaml
|
||||
callbacks:
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: examples/distill/MatrixGame2.0/validation.json
|
||||
sampling_steps: [4]
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- optical_flow.synthetic_optical_flow
|
||||
- common.fvd
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
# `synthetic_optical_flow` also needs actions in the validation
|
||||
# manifest (`action_path`) plus a repo-local calibration file.
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
```
|
||||
|
||||
Metric summaries are written under
|
||||
`<output_dir>/eval/step_<step>/inference_steps_<n>_rank_<rank>.json` and
|
||||
scalar means are logged with the `metrics/validation/...` prefix.
|
||||
Install `fastvideo[eval]` for optical-flow dependencies such as
|
||||
`ptlflow`; VBench metrics use the pinned submodule under
|
||||
`fastvideo/third_party/eval/vbench`. FVD uses `ref_video` entries from
|
||||
the validation manifest when present, or the standard
|
||||
`FASTVIDEO_FVD_REF_FEATURES` / eval-cache reference feature path.
|
||||
|
||||
### CLI
|
||||
|
||||
```bash
|
||||
@@ -350,5 +388,3 @@ sweep several baselines into a table, see
|
||||
|
||||
- **MIND** metrics. Depend on a separate `vipe` upstream submodule.
|
||||
- **VBench-2.0**. Sibling vbench2 package; needs its own port.
|
||||
- **Training-time eval callback** (`EvalCallback`) and the
|
||||
`RolloutEvaluator` helper.
|
||||
|
||||
@@ -63,7 +63,7 @@ class ThirdPersonCalibration:
|
||||
beta_fwd: float
|
||||
beta_strafe: float
|
||||
focal_length: float
|
||||
r_z: float
|
||||
r_z: float = 0.0
|
||||
r_y: float = 0.0
|
||||
init_pitch: float = 0.0
|
||||
notes: str = ""
|
||||
|
||||
@@ -261,6 +261,22 @@ def _mm_fp4(
|
||||
)
|
||||
|
||||
|
||||
def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor:
|
||||
"""Coerce an activation to a dtype the FP4 linear accepts.
|
||||
|
||||
The pre-attention norm can emit fp32 (e.g. in eager mode, without the
|
||||
torch.compile fusion that keeps it bf16). The FP4 linear emits bf16
|
||||
regardless (see _mm_fp4 out dtype), so cast fp32 -> bf16 rather than
|
||||
failing, matching the sibling fastvideo/layers/fp4linear.py. Non-floating
|
||||
inputs (e.g. int/bool) are a genuine error and are rejected fast.
|
||||
"""
|
||||
if not x.is_floating_point():
|
||||
raise TypeError(f"fp4 linear expects floating-point inputs, got {x.dtype}")
|
||||
if x.dtype not in (torch.bfloat16, torch.float16):
|
||||
x = x.to(torch.bfloat16)
|
||||
return x
|
||||
|
||||
|
||||
class NVFP4QuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def __init__(self, layer_prefix: str = ""):
|
||||
@@ -285,8 +301,7 @@ class NVFP4QuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
SfLayout, _, _ = _require_flashinfer()
|
||||
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, (
|
||||
f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
|
||||
x = _coerce_fp4_input_dtype(x)
|
||||
x_2d = x.view(-1, x.shape[-1])
|
||||
x_fp4, x_scale = _nvfp4_quantize(
|
||||
x_2d,
|
||||
@@ -332,8 +347,7 @@ class NVFP4QuantizeMethod(QuantizeMethodBase):
|
||||
if x_scale.dim() > 2:
|
||||
x_scale = x_scale.view(-1, x_scale.shape[-1])
|
||||
else:
|
||||
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, (
|
||||
f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
|
||||
x = _coerce_fp4_input_dtype(x)
|
||||
x = x.view(-1, x.shape[-1])
|
||||
x_global_sf = self.x_global_sf
|
||||
x_fp4, x_scale = _nvfp4_quantize(
|
||||
|
||||
@@ -295,6 +295,7 @@ def get_1d_rotary_pos_embed(
|
||||
interpolation_factor: float = 1.0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
use_real: bool = True,
|
||||
freqs_dtype: torch.dtype | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
@@ -319,12 +320,16 @@ def get_1d_rotary_pos_embed(
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos).float()
|
||||
|
||||
# freqs_dtype is an alias for dtype (Diffusers-compatible calling convention).
|
||||
if freqs_dtype is not None:
|
||||
dtype = freqs_dtype
|
||||
|
||||
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
||||
# has some connection to NTK literature
|
||||
if theta_rescale_factor != 1.0:
|
||||
theta *= theta_rescale_factor**(dim / (dim - 2))
|
||||
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(dtype) / dim)) # [D/2]
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2, device=pos.device)[:(dim // 2)].to(dtype) / dim)) # [D/2]
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
freqs_cos = freqs.cos() # [S, D/2]
|
||||
freqs_sin = freqs.sin() # [S, D/2]
|
||||
@@ -445,6 +450,21 @@ def get_nd_rotary_pos_embed(
|
||||
return cos, sin
|
||||
|
||||
|
||||
_ROTARY_POS_EMBED_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
|
||||
# Bound the table cache so long-running servers / causal models (which vary
|
||||
# start_frame per frame) cannot grow it without limit; entries are large float64
|
||||
# tensors. Least-recently-used eviction keeps the active resolution(s) hot while
|
||||
# capping memory.
|
||||
_ROTARY_POS_EMBED_CACHE_MAXSIZE = 16
|
||||
|
||||
|
||||
def _hashable(value: Any) -> Any:
|
||||
"""Return a hashable view of a scalar or sequence for use in a cache key."""
|
||||
if isinstance(value, list | tuple):
|
||||
return tuple(value)
|
||||
return value
|
||||
|
||||
|
||||
def get_rotary_pos_embed(
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
@@ -495,6 +515,31 @@ def get_rotary_pos_embed(
|
||||
sp_rank = 0
|
||||
sp_world_size = 1
|
||||
|
||||
# Memoize on every output-affecting argument; the table is constant across
|
||||
# denoising steps, so this avoids recomputing the float64 cos/sin tables.
|
||||
cache_key = (
|
||||
_hashable(rope_sizes),
|
||||
tuple(rope_dim_list),
|
||||
rope_theta,
|
||||
_hashable(theta_rescale_factor),
|
||||
_hashable(interpolation_factor),
|
||||
shard_dim,
|
||||
sp_rank,
|
||||
sp_world_size,
|
||||
dtype,
|
||||
start_frame,
|
||||
use_real,
|
||||
)
|
||||
cached = _ROTARY_POS_EMBED_CACHE.get(cache_key)
|
||||
if cached is not None:
|
||||
# Move to most-recently-used position so the active table is not evicted
|
||||
# when several resolutions / buckets share the process (LRU recency).
|
||||
# Pop with a default: a concurrent eviction between the get() above and
|
||||
# here would otherwise raise KeyError on the hit path.
|
||||
if _ROTARY_POS_EMBED_CACHE.pop(cache_key, None) is not None:
|
||||
_ROTARY_POS_EMBED_CACHE[cache_key] = cached
|
||||
return cached
|
||||
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
@@ -508,6 +553,14 @@ def get_rotary_pos_embed(
|
||||
start_frame=start_frame,
|
||||
use_real=use_real,
|
||||
)
|
||||
# The returned tensors are shared cache entries: callers must never mutate
|
||||
# them in place. Note .to(device) is an identity alias when the tensor is
|
||||
# already on the target device (e.g. CPU runs), so it does NOT guarantee a
|
||||
# copy — treat the tables as read-only and copy before any in-place op.
|
||||
# Reached only on a miss, so evict the least-recently-used entry at capacity.
|
||||
if len(_ROTARY_POS_EMBED_CACHE) >= _ROTARY_POS_EMBED_CACHE_MAXSIZE:
|
||||
_ROTARY_POS_EMBED_CACHE.pop(next(iter(_ROTARY_POS_EMBED_CACHE)))
|
||||
_ROTARY_POS_EMBED_CACHE[cache_key] = (freqs_cos, freqs_sin)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
|
||||
@@ -323,9 +323,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -437,6 +437,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.teacher_forcing_block_mask = None
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
assert self.num_frame_per_block <= 3
|
||||
self.independent_first_frame = False
|
||||
@@ -500,6 +501,70 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
return block_mask
|
||||
|
||||
@staticmethod
|
||||
def _prepare_teacher_forcing_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
|
||||
) -> BlockMask:
|
||||
"""Attention mask for the teacher-forcing ``[clean | noisy]`` sequence.
|
||||
|
||||
A noisy token attends to its own block plus the clean context of all
|
||||
strictly previous blocks; clean tokens are block-wise causal.
|
||||
"""
|
||||
if local_attn_size != -1:
|
||||
raise NotImplementedError(
|
||||
f"Teacher forcing ignores local_attn_size={local_attn_size}: "
|
||||
"unlike the block-wise causal mask, this mask always attends "
|
||||
"to the full clean context. Windowed teacher forcing is not "
|
||||
"implemented; use local_attn_size=-1 for teacher-forcing "
|
||||
"training.")
|
||||
total_length = num_frames * frame_seqlen * 2
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
|
||||
clean_ends = num_frames * frame_seqlen
|
||||
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_noise_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_noise_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
|
||||
attention_block_size = frame_seqlen * num_frame_per_block
|
||||
frame_indices = torch.arange(
|
||||
start=0, end=num_frames * frame_seqlen,
|
||||
step=attention_block_size, device=device, dtype=torch.long
|
||||
)
|
||||
for start in frame_indices:
|
||||
context_ends[start:start + attention_block_size] = start + attention_block_size
|
||||
|
||||
noisy_image_start_list = torch.arange(
|
||||
num_frames * frame_seqlen, total_length,
|
||||
step=attention_block_size, device=device, dtype=torch.long
|
||||
)
|
||||
noisy_image_end_list = noisy_image_start_list + attention_block_size
|
||||
for block_index, (start, end) in enumerate(zip(noisy_image_start_list, noisy_image_end_list)):
|
||||
noise_noise_starts[start:end] = start
|
||||
noise_noise_ends[start:end] = end
|
||||
noise_context_ends[start:end] = block_index * attention_block_size
|
||||
|
||||
def attention_mask(b, h, q_idx, kv_idx):
|
||||
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
|
||||
c1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
|
||||
c2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
|
||||
noise_mask = (q_idx >= clean_ends) & (c1 | c2)
|
||||
eye_mask = q_idx == kv_idx
|
||||
return eye_mask | clean_mask | noise_mask
|
||||
|
||||
block_mask = create_block_mask(
|
||||
attention_mask, B=None, H=None,
|
||||
Q_LEN=total_length + padded_length, KV_LEN=total_length + padded_length,
|
||||
_compile=False, device=device)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(f" cache a teacher-forcing mask with block size of {num_frame_per_block} frames")
|
||||
print(block_mask)
|
||||
|
||||
return block_mask
|
||||
|
||||
def _forward_inference(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -600,7 +665,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
else:
|
||||
causal_kwargs = {
|
||||
"kv_cache": kv_cache[block_index],
|
||||
"crossattn_cache": crossattn_cache[block_index],
|
||||
"crossattn_cache": (crossattn_cache[block_index]
|
||||
if crossattn_cache is not None else None),
|
||||
"current_start": current_start,
|
||||
"cache_start": cache_start,
|
||||
"block_mask": self.block_mask,
|
||||
@@ -628,9 +694,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
start_frame: int = 0,
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
teacher_forcing = clean_x is not None
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
@@ -663,15 +732,26 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.block_mask is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size
|
||||
)
|
||||
if teacher_forcing:
|
||||
if self.teacher_forcing_block_mask is None:
|
||||
self.teacher_forcing_block_mask = self._prepare_teacher_forcing_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
)
|
||||
block_mask = self.teacher_forcing_block_mask
|
||||
else:
|
||||
if self.block_mask is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size
|
||||
)
|
||||
block_mask = self.block_mask
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
@@ -679,6 +759,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
encoder_hidden_states_text = encoder_hidden_states
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
@@ -694,18 +775,35 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
if teacher_forcing:
|
||||
# Tile RoPE/modulation so clean frame i and noisy frame i share a position.
|
||||
clean_tokens = self.patch_embedding(clean_x).flatten(2).transpose(1, 2)
|
||||
hidden_states = torch.cat([clean_tokens, hidden_states], dim=1)
|
||||
if aug_t is None:
|
||||
aug_t = torch.zeros_like(timestep)
|
||||
_, timestep_proj_clean, _, _ = self.condition_embedder(
|
||||
aug_t.flatten(), encoder_hidden_states_text, None)
|
||||
timestep_proj_clean = timestep_proj_clean.unflatten(
|
||||
1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
timestep_proj = torch.cat([timestep_proj_clean, timestep_proj], dim=1)
|
||||
freqs_cis = (torch.cat([freqs_cos, freqs_cos], dim=0),
|
||||
torch.cat([freqs_sin, freqs_sin], dim=0))
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
block_mask=block_mask)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
block_mask=block_mask)
|
||||
|
||||
if teacher_forcing:
|
||||
hidden_states = hidden_states[:, hidden_states.shape[1] // 2:]
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user