Compare commits
53
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8f02827780 | ||
|
|
21047a6850 | ||
|
|
5ef40eff4c | ||
|
|
944637d363 | ||
|
|
29b66fe563 | ||
|
|
715d6a1d1a | ||
|
|
eb305db9d8 | ||
|
|
873efad60a | ||
|
|
7f86e8ab79 | ||
|
|
f772ee05bc | ||
|
|
0fc096d259 | ||
|
|
f3dea45bdc | ||
|
|
1666d1a66c | ||
|
|
e6dcbb6cb8 | ||
|
|
df7282432e | ||
|
|
b227e0e9b5 | ||
|
|
96426e9325 | ||
|
|
0c4cac7119 | ||
|
|
0098efa85e | ||
|
|
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",
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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,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,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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -1,15 +1,39 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
|
||||
# kernel package's FASTVIDEO_VSA_CUTEDSL: 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. Below sm90 a capability
|
||||
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
|
||||
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
|
||||
# JIT, observed on sm_89).
|
||||
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"
|
||||
except ImportError:
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
@@ -21,6 +45,12 @@ 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
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
@@ -87,9 +117,10 @@ if fa_version in ("2", "3"):
|
||||
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.
|
||||
# FA4 path: `flash_attn_func` (from `flash_attn_cute`) goes through a
|
||||
# registered torch.library custom op (with an FA4 backward on sm90+;
|
||||
# grad-enabled and GQA calls below sm90 route to FA2), 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:
|
||||
@@ -99,17 +130,6 @@ else:
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
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 +291,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)
|
||||
@@ -322,7 +338,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
|
||||
@@ -19,7 +20,7 @@ class SDPABackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SDPA"
|
||||
return "TORCH_SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SDPAImpl"]:
|
||||
@@ -49,9 +50,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 +125,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,
|
||||
|
||||
@@ -21,24 +21,35 @@ 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, )
|
||||
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 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_cute
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
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_flash
|
||||
return flash_attn_varlen_func_flash
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
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.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
@@ -13,7 +15,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"
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,511 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldConfig
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather_with_unpad,
|
||||
sequence_model_parallel_shard)
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
|
||||
from fastvideo.models.dits.wanvideo import (LayerNormScaleShift,
|
||||
PatchEmbed,
|
||||
WanTimeTextImageEmbedding,
|
||||
WanTransformer3DModel,
|
||||
WanTransformerBlock)
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
|
||||
|
||||
def _dreamx_invert_se3(transforms: torch.Tensor) -> torch.Tensor:
|
||||
assert transforms.shape[-2:] == (4, 4)
|
||||
rot_inv = transforms[..., :3, :3].transpose(-1, -2)
|
||||
out = torch.zeros_like(transforms)
|
||||
out[..., :3, :3] = rot_inv
|
||||
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", rot_inv,
|
||||
transforms[..., :3, 3])
|
||||
out[..., 3, 3] = 1.0
|
||||
return out.to(dtype=transforms.dtype)
|
||||
|
||||
|
||||
def _dreamx_lift_k(intrinsics: torch.Tensor) -> torch.Tensor:
|
||||
assert intrinsics.shape[-2:] == (3, 3)
|
||||
out = torch.zeros(intrinsics.shape[:-2] + (4, 4),
|
||||
device=intrinsics.device,
|
||||
dtype=intrinsics.dtype)
|
||||
out[..., :3, :3] = intrinsics
|
||||
out[..., 3, 3] = 1.0
|
||||
return out
|
||||
|
||||
|
||||
def _dreamx_invert_k(intrinsics: torch.Tensor) -> torch.Tensor:
|
||||
assert intrinsics.shape[-2:] == (3, 3)
|
||||
out = torch.zeros_like(intrinsics)
|
||||
out[..., 0, 0] = 1.0 / intrinsics[..., 0, 0]
|
||||
out[..., 1, 1] = 1.0 / intrinsics[..., 1, 1]
|
||||
out[..., 0, 2] = -intrinsics[..., 0, 2] / intrinsics[..., 0, 0]
|
||||
out[..., 1, 2] = -intrinsics[..., 1, 2] / intrinsics[..., 1, 1]
|
||||
out[..., 2, 2] = 1.0
|
||||
return out.to(dtype=intrinsics.dtype)
|
||||
|
||||
|
||||
def _dreamx_apply_tiled_projmat(feats: torch.Tensor,
|
||||
matrix: torch.Tensor) -> torch.Tensor:
|
||||
batch, num_heads, seq_len, feat_dim = feats.shape
|
||||
proj_dim = matrix.shape[-1]
|
||||
assert feat_dim % proj_dim == 0
|
||||
|
||||
if matrix.shape[1] == seq_len:
|
||||
feats = feats.view(batch, num_heads, seq_len, feat_dim // proj_dim,
|
||||
proj_dim)
|
||||
out = torch.einsum("btij,bntpj->bntpi", matrix, feats)
|
||||
return out.reshape(batch, num_heads, seq_len, feat_dim)
|
||||
|
||||
cameras = matrix.shape[1]
|
||||
assert seq_len > cameras and seq_len % cameras == 0
|
||||
feats = feats.reshape(batch, num_heads, cameras, -1,
|
||||
feat_dim // proj_dim, proj_dim)
|
||||
out = torch.einsum("bcij,bncpkj->bncpki", matrix, feats)
|
||||
return out.reshape(batch, num_heads, seq_len, feat_dim)
|
||||
|
||||
|
||||
def _dreamx_prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
viewmats: torch.Tensor, intrinsics: torch.Tensor):
|
||||
batch, num_heads, seq_len, head_dim = q.shape
|
||||
cameras = viewmats.shape[1]
|
||||
assert q.shape == k.shape == v.shape
|
||||
assert viewmats.shape == (batch, cameras, 4, 4)
|
||||
assert intrinsics.shape == (batch, cameras, 3, 3)
|
||||
assert head_dim % 4 == 0
|
||||
|
||||
intrinsics_norm = torch.zeros_like(intrinsics)
|
||||
intrinsics_norm[..., 0, 0] = intrinsics[..., 0, 0]
|
||||
intrinsics_norm[..., 1, 1] = intrinsics[..., 1, 1]
|
||||
intrinsics_norm[..., 2, 2] = 1.0
|
||||
|
||||
proj = torch.einsum("...ij,...jk->...ik",
|
||||
_dreamx_lift_k(intrinsics_norm), viewmats)
|
||||
proj_t = proj.transpose(-1, -2).to(dtype=viewmats.dtype)
|
||||
proj_inv = torch.einsum(
|
||||
"...ij,...jk->...ik",
|
||||
_dreamx_invert_se3(viewmats),
|
||||
_dreamx_lift_k(_dreamx_invert_k(intrinsics_norm)),
|
||||
).to(dtype=viewmats.dtype)
|
||||
|
||||
q = _dreamx_apply_tiled_projmat(q, proj_t)
|
||||
k = _dreamx_apply_tiled_projmat(k, proj_inv)
|
||||
v = _dreamx_apply_tiled_projmat(v, proj_inv)
|
||||
return q, k, v, proj
|
||||
|
||||
|
||||
class DreamXPropeSelfAttention(nn.Module):
|
||||
"""DreamX-World parallel PRoPE camera self-attention branch."""
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
attn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str | bool = True,
|
||||
eps: float = 1e-6,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
assert attn_dim % num_heads == 0
|
||||
self.attn_dim = attn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = attn_dim // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
|
||||
self.q_proj = ReplicatedLinear(dim,
|
||||
attn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.q_proj")
|
||||
self.k_proj = ReplicatedLinear(dim,
|
||||
attn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.k_proj")
|
||||
self.v_proj = ReplicatedLinear(dim,
|
||||
attn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.v_proj")
|
||||
self.out_proj = ReplicatedLinear(attn_dim,
|
||||
dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.out_proj")
|
||||
|
||||
if qk_norm == "rms_norm":
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps)
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps)
|
||||
elif qk_norm in (True, "rms_norm_across_heads"):
|
||||
self.norm_q = RMSNorm(attn_dim, eps=eps)
|
||||
self.norm_k = RMSNorm(attn_dim, eps=eps)
|
||||
elif qk_norm is False:
|
||||
self.norm_q = nn.Identity()
|
||||
self.norm_k = nn.Identity()
|
||||
else:
|
||||
raise ValueError(f"Unsupported qk_norm for DreamX PRoPE: {qk_norm}")
|
||||
|
||||
nn.init.zeros_(self.out_proj.weight)
|
||||
if self.out_proj.bias is not None:
|
||||
nn.init.zeros_(self.out_proj.bias)
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor,
|
||||
y_camera: dict[str, torch.Tensor]) -> torch.Tensor:
|
||||
if get_sp_world_size() > 1:
|
||||
# The transformer shards the sequence before the block loop and
|
||||
# this branch uses LocalAttention (no all-to-all): under
|
||||
# sequence parallelism each rank would attend only within its
|
||||
# own shard — silently wrong output. Fail loudly until this
|
||||
# path is ported to DistributedAttention and validated.
|
||||
raise NotImplementedError(
|
||||
"DreamXPropeSelfAttention does not support sequence "
|
||||
"parallelism yet (LocalAttention on a sharded sequence "
|
||||
"corrupts output). Run with sp_size=1.")
|
||||
batch_size, seq_len, _ = hidden_states.shape
|
||||
|
||||
query, _ = self.q_proj(hidden_states)
|
||||
key, _ = self.k_proj(hidden_states)
|
||||
value, _ = self.v_proj(hidden_states)
|
||||
|
||||
if self.qk_norm == "rms_norm":
|
||||
query = query.view(batch_size, seq_len, self.num_heads,
|
||||
self.head_dim)
|
||||
key = key.view(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
else:
|
||||
query = self.norm_q(query).view(batch_size, seq_len,
|
||||
self.num_heads, self.head_dim)
|
||||
key = self.norm_k(key).view(batch_size, seq_len, self.num_heads,
|
||||
self.head_dim)
|
||||
|
||||
value = value.view(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
query, key, value, output_projection = _dreamx_prope_qkv(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
viewmats=y_camera["viewmats"],
|
||||
intrinsics=y_camera["K"],
|
||||
)
|
||||
|
||||
out = self.attn(query.transpose(1, 2), key.transpose(1, 2),
|
||||
value.transpose(1, 2))
|
||||
out = _dreamx_apply_tiled_projmat(out.transpose(1, 2),
|
||||
output_projection).transpose(1, 2)
|
||||
out = out.flatten(2)
|
||||
out, _ = self.out_proj(out)
|
||||
return out
|
||||
|
||||
|
||||
class DreamXWorldTransformerBlock(WanTransformerBlock):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: int | None = None,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
add_control_adapter: bool = True,
|
||||
cam_method: str | None = "prope",
|
||||
attn_compress: int = 1,
|
||||
cam_self_attn_layers: tuple[int, ...] | None = None,
|
||||
layer_idx: int | None = None):
|
||||
super().__init__(dim, ffn_dim, num_heads, qk_norm, cross_attn_norm,
|
||||
eps, added_kv_proj_dim,
|
||||
supported_attention_backends, quant_config, prefix)
|
||||
self.cam_self_attn = None
|
||||
add_cam_attn = add_control_adapter and cam_method == "prope"
|
||||
if add_cam_attn and cam_self_attn_layers is not None:
|
||||
add_cam_attn = layer_idx in cam_self_attn_layers
|
||||
if add_cam_attn:
|
||||
if num_heads % attn_compress != 0 or dim % attn_compress != 0:
|
||||
raise ValueError("DreamX attn_compress must divide dim and num_heads")
|
||||
self.cam_self_attn = DreamXPropeSelfAttention(
|
||||
dim,
|
||||
dim // attn_compress,
|
||||
num_heads // attn_compress,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.cam_self_attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
original_seq_len: int,
|
||||
y_camera: dict[str, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
orig_dtype = hidden_states.dtype
|
||||
|
||||
if temb.dim() == 4:
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb.float()).chunk(
|
||||
6, dim=2)
|
||||
shift_msa = shift_msa.squeeze(2)
|
||||
scale_msa = scale_msa.squeeze(2)
|
||||
gate_msa = gate_msa.squeeze(2)
|
||||
c_shift_msa = c_shift_msa.squeeze(2)
|
||||
c_scale_msa = c_scale_msa.squeeze(2)
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
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))
|
||||
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
|
||||
attn_output, _ = self.attn1(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
original_seq_len,
|
||||
freqs_cis=freqs_cis,
|
||||
)
|
||||
attn_output = attn_output.flatten(2)
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
if self.cam_self_attn is not None and y_camera is not None:
|
||||
attn_output = attn_output + self.cam_self_attn(
|
||||
norm_hidden_states, y_camera)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DreamXWorldTransformer3DModel(WanTransformer3DModel):
|
||||
_fsdp_shard_conditions = DreamXWorldConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = DreamXWorldConfig()._compile_conditions
|
||||
_supported_attention_backends = DreamXWorldConfig(
|
||||
)._supported_attention_backends
|
||||
param_names_mapping = DreamXWorldConfig().param_names_mapping
|
||||
reverse_param_names_mapping = DreamXWorldConfig().reverse_param_names_mapping
|
||||
lora_param_names_mapping = DreamXWorldConfig().lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: DreamXWorldConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
BaseDiT.__init__(self, config=config, hf_config=hf_config)
|
||||
self.quant_config = config.quant_config
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.text_len = config.text_len
|
||||
|
||||
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
|
||||
|
||||
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
|
||||
embed_dim=inner_dim,
|
||||
patch_size=config.patch_size,
|
||||
flatten=False)
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=config.freq_dim,
|
||||
text_embed_dim=config.text_dim,
|
||||
image_embed_dim=config.image_dim,
|
||||
)
|
||||
self.blocks = nn.ModuleList([
|
||||
DreamXWorldTransformerBlock(
|
||||
inner_dim,
|
||||
config.ffn_dim,
|
||||
config.num_attention_heads,
|
||||
config.qk_norm,
|
||||
config.cross_attn_norm,
|
||||
config.eps,
|
||||
config.added_kv_proj_dim,
|
||||
self._supported_attention_backends,
|
||||
quant_config=config.quant_config,
|
||||
prefix=f"{config.prefix}.blocks.{i}",
|
||||
add_control_adapter=config.add_control_adapter,
|
||||
cam_method=config.cam_method,
|
||||
attn_compress=config.attn_compress,
|
||||
cam_self_attn_layers=config.cam_self_attn_layers,
|
||||
layer_idx=i)
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
self.norm_out = LayerNormScaleShift(inner_dim,
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
self.gradient_checkpointing = False
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
guidance=None,
|
||||
y_camera: dict[str, torch.Tensor] | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
orig_dtype = hidden_states.dtype
|
||||
if encoder_hidden_states is not None and not isinstance(
|
||||
encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
list) and len(encoder_hidden_states_image) > 0:
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
else:
|
||||
encoder_hidden_states_image = None
|
||||
|
||||
batch_size, _, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
d = self.hidden_size // self.num_attention_heads
|
||||
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||
(post_patch_num_frames, post_patch_height, post_patch_width),
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
|
||||
rope_theta=10000)
|
||||
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
|
||||
freqs_sin.to(hidden_states.device).float())
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
hidden_states, original_seq_len = sequence_model_parallel_shard(
|
||||
hidden_states, dim=1)
|
||||
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
timestep = timestep.flatten()
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
encoder_hidden_states_image,
|
||||
timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
timestep_proj = timestep_proj.unflatten(2, (6, -1))
|
||||
else:
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
else:
|
||||
encoder_hidden_states = encoder_hidden_states_image
|
||||
|
||||
if current_platform.is_mps() or current_platform.is_npu():
|
||||
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
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, original_seq_len, y_camera)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, original_seq_len,
|
||||
y_camera=y_camera)
|
||||
|
||||
if temb.dim() == 3:
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(0) +
|
||||
temb.unsqueeze(2)).chunk(2, dim=2)
|
||||
shift = shift.squeeze(2)
|
||||
scale = scale.squeeze(2)
|
||||
else:
|
||||
shift, scale = (self.scale_shift_table +
|
||||
temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = sequence_model_parallel_all_gather_with_unpad(
|
||||
hidden_states, original_seq_len, dim=1)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
EntryClass = DreamXWorldTransformer3DModel
|
||||
@@ -0,0 +1,920 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World autoregressive causal DiT.
|
||||
|
||||
Adapted from DreamX-World's Apache-2.0
|
||||
``wan/modules/causal_camera_model_2_2_prope_infinity.py``. The implementation is
|
||||
kept native to FastVideo: no production import from DreamX, Diffusers, or
|
||||
Transformers is required.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.dreamx_world import (_dreamx_apply_tiled_projmat,
|
||||
_dreamx_prope_qkv)
|
||||
|
||||
|
||||
def attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
|
||||
# Deliberately raw SDPA rather than fastvideo.attention.LocalAttention:
|
||||
# (1) LocalAttention dispatches through the attention-backend registry, so
|
||||
# FLASH_ATTN could be selected and its kernel is not bit-identical to
|
||||
# torch SDPA — the AR KV-cache rollout must stay numerically frozen;
|
||||
# (2) LocalAttention requires an active ForwardContext, which direct
|
||||
# transformer invocations (parity tests) do not set;
|
||||
# (3) the sibling causal model keeps raw SDPA in the same KV-cache window
|
||||
# path (matrixgame2/causal_model.py).
|
||||
# Sequence-parallel gap: this model never shards the sequence; run with
|
||||
# sp_size=1 (see fastvideo/layers/AGENTS.md on documenting raw SDPA).
|
||||
q_bhld = q.transpose(1, 2)
|
||||
k_bhld = k.transpose(1, 2)
|
||||
v_bhld = v.transpose(1, 2)
|
||||
out = F.scaled_dot_product_attention(q_bhld, k_bhld, v_bhld, dropout_p=0.0)
|
||||
return out.transpose(1, 2)
|
||||
|
||||
|
||||
def prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
viewmats: torch.Tensor, Ks: torch.Tensor):
|
||||
q, k, v, output_projection = _dreamx_prope_qkv(q, k, v, viewmats, Ks)
|
||||
|
||||
def apply_fn_o(x: torch.Tensor) -> torch.Tensor:
|
||||
return _dreamx_apply_tiled_projmat(x, output_projection)
|
||||
|
||||
return q, k, v, apply_fn_o
|
||||
|
||||
|
||||
def sinusoidal_embedding_1d(dim, position):
|
||||
assert dim % 2 == 0
|
||||
half = dim // 2
|
||||
position = position.type(torch.float64)
|
||||
sinusoid = torch.outer(
|
||||
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
|
||||
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
|
||||
|
||||
def rope_params(max_seq_len, dim, theta=10000):
|
||||
assert dim % 2 == 0
|
||||
freqs = torch.outer(
|
||||
torch.arange(max_seq_len),
|
||||
1.0 / torch.pow(theta,
|
||||
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
|
||||
return torch.polar(torch.ones_like(freqs), freqs)
|
||||
|
||||
|
||||
class WanRMSNorm(nn.Module):
|
||||
"""Kept private instead of fastvideo.layers.layernorm.RMSNorm.
|
||||
|
||||
The official DreamX-World ``model_2_2.py`` computes the RMS statistics in
|
||||
the *input* dtype — the upstream code has the fp32 upcast explicitly
|
||||
commented out (``# return self._norm(x.float())...``). FastVideo's RMSNorm
|
||||
always normalizes in fp32, which is not bit-identical under bf16, so the
|
||||
verbatim implementation stays.
|
||||
"""
|
||||
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return self._norm(x).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class WanLayerNorm(nn.LayerNorm):
|
||||
"""Kept private instead of fastvideo.layers.layernorm.FP32LayerNorm.
|
||||
|
||||
The official DreamX-World ``model_2_2.py`` normalizes in the *input* dtype
|
||||
(no ``x.float()`` upcast, unlike Wan2.1). FP32LayerNorm casts input and
|
||||
affine params to fp32, which is not bit-identical under bf16, so the
|
||||
verbatim implementation stays.
|
||||
"""
|
||||
|
||||
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||
|
||||
def forward(self, x):
|
||||
return super().forward(x).type_as(x)
|
||||
|
||||
|
||||
class WanCrossAttention(nn.Module):
|
||||
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.q = ReplicatedLinear(dim, dim)
|
||||
self.k = ReplicatedLinear(dim, dim)
|
||||
self.v = ReplicatedLinear(dim, dim)
|
||||
self.o = ReplicatedLinear(dim, dim)
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def _kv(self, context, b, n, d):
|
||||
k, _ = self.k(context)
|
||||
k = self.norm_k(k).view(b, -1, n, d)
|
||||
v, _ = self.v(context)
|
||||
v = v.view(b, -1, n, d)
|
||||
return k, v
|
||||
|
||||
def forward(self, x, context, context_lens, crossattn_cache=None):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
q, _ = self.q(x)
|
||||
q = self.norm_q(q).view(b, -1, n, d)
|
||||
|
||||
if crossattn_cache is not None:
|
||||
if not crossattn_cache["is_init"]:
|
||||
crossattn_cache["is_init"] = True
|
||||
k, v = self._kv(context, b, n, d)
|
||||
crossattn_cache["k"] = k
|
||||
crossattn_cache["v"] = v
|
||||
else:
|
||||
k = crossattn_cache["k"]
|
||||
v = crossattn_cache["v"]
|
||||
else:
|
||||
k, v = self._kv(context, b, n, d)
|
||||
|
||||
x = attention(q, k, v)
|
||||
x = x.flatten(2)
|
||||
out, _ = self.o(x)
|
||||
return out
|
||||
|
||||
|
||||
def block_relativistic_rope(x, grid_sizes, freqs, start_frame=0, relative_frame_indices=None):
|
||||
"""
|
||||
Apply Block-Relativistic RoPE to input tensor.
|
||||
Adapted from Infinity-RoPE (https://arxiv.org/abs/2511.20649).
|
||||
|
||||
Args:
|
||||
x: Input tensor [B, L, num_heads, head_dim]
|
||||
grid_sizes: Tensor [B, 3] containing (F, H, W)
|
||||
freqs: RoPE frequencies
|
||||
start_frame: Starting frame index for sequential RoPE
|
||||
relative_frame_indices: Optional tensor [F] specifying explicit frame indices
|
||||
for Block-Relativistic RoPE. Overrides start_frame if provided.
|
||||
"""
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
|
||||
seq_len, n, -1, 2))
|
||||
|
||||
if relative_frame_indices is not None:
|
||||
frame_indices = relative_frame_indices.long()
|
||||
freqs_temporal = freqs[0][frame_indices].view(f, 1, 1, -1).expand(f, h, w, -1)
|
||||
else:
|
||||
freqs_temporal = freqs[0][start_frame:start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1)
|
||||
|
||||
freqs_i = torch.cat([
|
||||
freqs_temporal,
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||
], dim=-1).reshape(seq_len, 1, -1)
|
||||
|
||||
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]])
|
||||
output.append(x_i)
|
||||
|
||||
return torch.stack(output).type_as(x)
|
||||
|
||||
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
"""Self-attention with KV cache and Block-Relativistic RoPE for causal inference."""
|
||||
|
||||
def __init__(self, dim, num_heads, local_attn_size=6, sink_size=1,
|
||||
qk_norm=True, eps=1e-6):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
|
||||
|
||||
self.q = ReplicatedLinear(dim, dim)
|
||||
self.k = ReplicatedLinear(dim, dim)
|
||||
self.v = ReplicatedLinear(dim, dim)
|
||||
self.o = ReplicatedLinear(dim, dim)
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs, kv_cache,
|
||||
current_start=0, cache_start=None, sink_recache_after_switch=False):
|
||||
"""
|
||||
Args:
|
||||
x: Shape [B, L, C]
|
||||
seq_lens: Shape [B]
|
||||
grid_sizes: Shape [B, 3] containing (F, H, W)
|
||||
freqs: RoPE frequencies [1024, head_dim / 2]
|
||||
kv_cache: Dict with 'k', 'v', 'global_end_index', 'local_end_index'
|
||||
current_start: Current position in the global token sequence
|
||||
"""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
q, _ = self.q(x)
|
||||
q = self.norm_q(q).view(b, s, n, d)
|
||||
k, _ = self.k(x)
|
||||
k = self.norm_k(k).view(b, s, n, d)
|
||||
v, _ = self.v(x)
|
||||
v = v.view(b, s, n, d)
|
||||
|
||||
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
|
||||
num_new_frames = grid_sizes[0][0].item()
|
||||
current_end = current_start + q.shape[1]
|
||||
sink_tokens = self.sink_size * frame_seqlen
|
||||
kv_cache_size = kv_cache["k"].shape[1]
|
||||
num_new_tokens = q.shape[1]
|
||||
|
||||
cache_update_info = None
|
||||
is_recompute = current_end <= kv_cache["global_end_index"].item() and current_start > 0
|
||||
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
# === ROLLING MODE: cache full, evict oldest non-sink tokens ===
|
||||
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
|
||||
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
temp_k = kv_cache["k"].detach().clone()
|
||||
temp_v = kv_cache["v"].detach().clone()
|
||||
temp_k[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
temp_k[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
temp_v[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
temp_v[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
|
||||
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
|
||||
|
||||
# Block-Relativistic RoPE: query uses window-relative indices
|
||||
query_relative_indices = torch.arange(
|
||||
self.local_attn_size - num_new_frames, self.local_attn_size, device=q.device)
|
||||
roped_query = block_relativistic_rope(
|
||||
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
|
||||
|
||||
# Block-Relativistic RoPE: cached K uses position-in-window indices
|
||||
num_cache_frames = local_end_index // frame_seqlen
|
||||
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
|
||||
cache_grid_sizes = grid_sizes.clone()
|
||||
cache_grid_sizes[0, 0] = num_cache_frames
|
||||
roped_temp_k = block_relativistic_rope(
|
||||
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
|
||||
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
|
||||
|
||||
cache_update_info = {
|
||||
"action": "roll_and_insert",
|
||||
"sink_tokens": sink_tokens,
|
||||
"num_rolled_tokens": num_rolled_tokens,
|
||||
"num_evicted_tokens": num_evicted_tokens,
|
||||
"local_start_index": local_start_index,
|
||||
"local_end_index": local_end_index,
|
||||
"write_start_index": write_start_index,
|
||||
"write_end_index": local_end_index,
|
||||
"new_k": k[:, roped_offset:roped_offset + write_len],
|
||||
"new_v": v[:, roped_offset:roped_offset + write_len],
|
||||
"current_end": current_end,
|
||||
"is_recompute": is_recompute
|
||||
}
|
||||
else:
|
||||
# === DIRECT INSERT MODE: cache not yet full ===
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
temp_k = kv_cache["k"].detach().clone()
|
||||
temp_v = kv_cache["v"].detach().clone()
|
||||
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
if sink_recache_after_switch:
|
||||
write_start_index = local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
|
||||
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
|
||||
|
||||
# RoPE with relative indices (growing sequentially before cache fills)
|
||||
current_frame_in_window = local_start_index // frame_seqlen
|
||||
query_relative_indices = torch.arange(
|
||||
current_frame_in_window, current_frame_in_window + num_new_frames, device=q.device)
|
||||
roped_query = block_relativistic_rope(
|
||||
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
|
||||
|
||||
num_cache_frames = local_end_index // frame_seqlen
|
||||
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
|
||||
cache_grid_sizes = grid_sizes.clone()
|
||||
cache_grid_sizes[0, 0] = num_cache_frames
|
||||
roped_temp_k = block_relativistic_rope(
|
||||
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
|
||||
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
|
||||
|
||||
cache_update_info = {
|
||||
"action": "direct_insert",
|
||||
"local_start_index": local_start_index,
|
||||
"local_end_index": local_end_index,
|
||||
"write_start_index": write_start_index,
|
||||
"write_end_index": local_end_index,
|
||||
"new_k": k[:, roped_offset:roped_offset + write_len],
|
||||
"new_v": v[:, roped_offset:roped_offset + write_len],
|
||||
"current_end": current_end,
|
||||
"is_recompute": is_recompute
|
||||
}
|
||||
|
||||
# Attention: sink tokens + local window
|
||||
if sink_tokens > 0:
|
||||
local_budget = self.max_attention_size - sink_tokens
|
||||
k_sink = roped_temp_k[:, :sink_tokens]
|
||||
v_sink = temp_v[:, :sink_tokens]
|
||||
if local_budget > 0:
|
||||
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
|
||||
k_local = roped_temp_k[:, local_start_for_window:local_end_index]
|
||||
v_local = temp_v[:, local_start_for_window:local_end_index]
|
||||
k_cat = torch.cat([k_sink, k_local], dim=1)
|
||||
v_cat = torch.cat([v_sink, v_local], dim=1)
|
||||
else:
|
||||
k_cat = k_sink
|
||||
v_cat = v_sink
|
||||
x = attention(roped_query, k_cat, v_cat)
|
||||
else:
|
||||
window_start = max(0, local_end_index - self.max_attention_size)
|
||||
x = attention(
|
||||
roped_query,
|
||||
roped_temp_k[:, window_start:local_end_index],
|
||||
temp_v[:, window_start:local_end_index])
|
||||
|
||||
x = x.flatten(2)
|
||||
x, _ = self.o(x)
|
||||
return x, (current_end, local_end_index, cache_update_info)
|
||||
|
||||
|
||||
class CausalPropeSelfAttention(nn.Module):
|
||||
"""PRoPE self-attention with optional KV cache for camera-controlled inference."""
|
||||
|
||||
def __init__(self, dim, attn_dim, num_heads, window_size=(-1, -1),
|
||||
local_attn_size=-1, sink_size=0, qk_norm=True, eps=1e-6):
|
||||
assert dim % num_heads == 0
|
||||
assert attn_dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.attn_dim = attn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = attn_dim // num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.window_size = window_size
|
||||
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
|
||||
|
||||
self.q_proj = ReplicatedLinear(dim, attn_dim)
|
||||
self.k_proj = ReplicatedLinear(dim, attn_dim)
|
||||
self.v_proj = ReplicatedLinear(dim, attn_dim)
|
||||
self.out_proj = ReplicatedLinear(attn_dim, dim)
|
||||
|
||||
self.norm_q = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
nn.init.zeros_(self.out_proj.weight)
|
||||
nn.init.zeros_(self.out_proj.bias)
|
||||
|
||||
def forward(self, x, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
|
||||
kv_cache=None, current_start=0, cache_start=None,
|
||||
sink_recache_after_switch=False, cache_update_policy="commit_detached"):
|
||||
"""
|
||||
Args:
|
||||
x: Shape [B, L, C]
|
||||
cam_viewmats: Camera view matrices
|
||||
cam_K: Camera intrinsics
|
||||
kv_cache: Optional KV cache dict. When None, runs full attention over current chunk.
|
||||
"""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
q, _ = self.q_proj(x)
|
||||
q = self.norm_q(q).view(b, s, n, d)
|
||||
k, _ = self.k_proj(x)
|
||||
k = self.norm_k(k).view(b, s, n, d)
|
||||
v, _ = self.v_proj(x)
|
||||
v = v.view(b, s, n, d)
|
||||
|
||||
# Apply PRoPE (Positional Rotary Position Embedding from camera parameters)
|
||||
q_t, k_t, v_t, apply_fn_o = prope_qkv(
|
||||
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2),
|
||||
viewmats=cam_viewmats, Ks=cam_K)
|
||||
proped_q = q_t.transpose(1, 2)
|
||||
proped_k = k_t.transpose(1, 2)
|
||||
proped_v = v_t.transpose(1, 2)
|
||||
|
||||
if kv_cache is None:
|
||||
# No cache: full attention over current chunk
|
||||
x_out = attention(proped_q, proped_k, proped_v)
|
||||
else:
|
||||
# KV cache mode with rolling cache support
|
||||
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
|
||||
num_new_tokens = s
|
||||
current_end = current_start + num_new_tokens
|
||||
sink_tokens = self.sink_size * frame_seqlen
|
||||
kv_cache_size = kv_cache["k"].shape[1]
|
||||
is_recompute = (current_end <= kv_cache["global_end_index"].item()) and (current_start > 0)
|
||||
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
# === ROLLING MODE ===
|
||||
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
|
||||
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
if cache_update_policy != "none":
|
||||
with torch.no_grad():
|
||||
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
with torch.no_grad():
|
||||
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
|
||||
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
|
||||
else:
|
||||
# === DIRECT INSERT MODE ===
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
if cache_update_policy != "none":
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
if sink_recache_after_switch:
|
||||
write_start_index = local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
with torch.no_grad():
|
||||
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
|
||||
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
|
||||
|
||||
# Attention: sink tokens + local window
|
||||
if sink_tokens > 0:
|
||||
local_budget = self.max_attention_size - sink_tokens
|
||||
k_sink = kv_cache["k"][:, :sink_tokens].detach()
|
||||
v_sink = kv_cache["v"][:, :sink_tokens].detach()
|
||||
if local_budget > 0:
|
||||
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
|
||||
k_local = kv_cache["k"][:, local_start_for_window:local_end_index].detach()
|
||||
v_local = kv_cache["v"][:, local_start_for_window:local_end_index].detach()
|
||||
k_cat = torch.cat([k_sink, k_local], dim=1)
|
||||
v_cat = torch.cat([v_sink, v_local], dim=1)
|
||||
else:
|
||||
k_cat = k_sink
|
||||
v_cat = v_sink
|
||||
x_out = attention(proped_q, k_cat, v_cat)
|
||||
else:
|
||||
window_start = max(0, local_end_index - self.max_attention_size)
|
||||
x_out = attention(
|
||||
proped_q,
|
||||
kv_cache["k"][:, window_start:local_end_index].detach(),
|
||||
kv_cache["v"][:, window_start:local_end_index].detach())
|
||||
|
||||
if not is_recompute and cache_update_policy != "none":
|
||||
kv_cache["global_end_index"].fill_(current_end)
|
||||
kv_cache["local_end_index"].fill_(local_end_index)
|
||||
|
||||
# Apply inverse PRoPE
|
||||
x = apply_fn_o(x_out.transpose(1, 2)).transpose(1, 2)
|
||||
x = x.flatten(2)
|
||||
x, _ = self.out_proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class CausalWanAttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self, dim, ffn_dim, num_heads, local_attn_size=-1, sink_size=0,
|
||||
qk_norm=True, cross_attn_norm=False, eps=1e-6, **kwargs):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
self.add_control_adapter = kwargs.get('add_control_adapter', False)
|
||||
self.cam_method = kwargs.get('cam_method')
|
||||
self.attn_compress = kwargs.get('attn_compress', 1)
|
||||
self.layer_idx = kwargs.get('layer_idx')
|
||||
cam_self_attn_layers = kwargs.get('cam_self_attn_layers')
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(dim, eps)
|
||||
self.self_attn = CausalWanSelfAttention(
|
||||
dim, num_heads, local_attn_size, sink_size, qk_norm, eps)
|
||||
self.norm3 = WanLayerNorm(
|
||||
dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.cross_attn = WanCrossAttention(dim, num_heads, (-1, -1), qk_norm, eps)
|
||||
self.norm2 = WanLayerNorm(dim, eps)
|
||||
# nn.Linear (not ReplicatedLinear) on purpose: the official checkpoint
|
||||
# stores these as positional Sequential keys (ffn.0 / ffn.2) that the
|
||||
# copy-only converter and the strict-load tests require verbatim, and
|
||||
# ReplicatedLinear's (out, bias) tuple return cannot compose inside
|
||||
# nn.Sequential without renaming the state-dict surface.
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn_dim, dim))
|
||||
|
||||
# PRoPE self-attention branch for camera control
|
||||
add_cam_attn = self.add_control_adapter and self.cam_method == 'prope'
|
||||
if add_cam_attn and cam_self_attn_layers is not None:
|
||||
add_cam_attn = self.layer_idx in cam_self_attn_layers
|
||||
if add_cam_attn:
|
||||
self.cam_self_attn = CausalPropeSelfAttention(
|
||||
dim, dim // self.attn_compress, num_heads,
|
||||
local_attn_size=local_attn_size, sink_size=sink_size,
|
||||
qk_norm=qk_norm, eps=eps)
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(self, x, e, seq_lens, grid_sizes, freqs, context, context_lens,
|
||||
kv_cache, crossattn_cache=None, current_start=0, cache_start=None,
|
||||
cam_viewmats=None, cam_K=None, sink_recache_after_switch=False,
|
||||
cache_update_policy="commit_detached"):
|
||||
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
|
||||
e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)
|
||||
|
||||
# self-attention
|
||||
attn_input = (self.norm1(x).unflatten(
|
||||
dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2)
|
||||
y, cache_update_info = self.self_attn(
|
||||
attn_input, seq_lens, grid_sizes, freqs, kv_cache,
|
||||
current_start, cache_start, sink_recache_after_switch)
|
||||
|
||||
# PRoPE camera attention (parallel branch)
|
||||
if hasattr(self, 'cam_self_attn') and cam_viewmats is not None and cam_K is not None:
|
||||
prope_kv_cache = None
|
||||
if kv_cache is not None and "prope_k" in kv_cache:
|
||||
prope_kv_cache = {
|
||||
"k": kv_cache["prope_k"],
|
||||
"v": kv_cache["prope_v"],
|
||||
"global_end_index": kv_cache["prope_global_end_index"],
|
||||
"local_end_index": kv_cache["prope_local_end_index"],
|
||||
}
|
||||
y = y + self.cam_self_attn(
|
||||
attn_input, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
|
||||
kv_cache=prope_kv_cache, current_start=current_start,
|
||||
cache_start=cache_start, cache_update_policy=cache_update_policy)
|
||||
|
||||
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(1, 2)
|
||||
|
||||
# cross-attention & FFN
|
||||
x = x + self.cross_attn(self.norm3(x), context, context_lens,
|
||||
crossattn_cache=crossattn_cache)
|
||||
y = self.ffn(
|
||||
(self.norm2(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
|
||||
* (1 + e[4]) + e[3]).flatten(1, 2))
|
||||
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[5]).flatten(1, 2)
|
||||
|
||||
return x, cache_update_info
|
||||
|
||||
|
||||
class CausalHead(nn.Module):
|
||||
|
||||
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.out_dim = out_dim
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
out_dim = math.prod(patch_size) * out_dim
|
||||
self.norm = WanLayerNorm(dim, eps)
|
||||
self.head = ReplicatedLinear(dim, out_dim)
|
||||
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
||||
|
||||
def forward(self, x, e):
|
||||
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
|
||||
e = (self.modulation.unsqueeze(1) + e).chunk(2, dim=2)
|
||||
x, _ = self.head(
|
||||
self.norm(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
|
||||
* (1 + e[1]) + e[0])
|
||||
return x
|
||||
|
||||
|
||||
class DreamXWorldARTransformer3DModel(BaseDiT):
|
||||
"""DreamX-World-5B autoregressive causal transformer."""
|
||||
|
||||
_fsdp_shard_conditions = DreamXWorldARConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = DreamXWorldARConfig()._compile_conditions
|
||||
_supported_attention_backends = DreamXWorldARConfig()._supported_attention_backends
|
||||
param_names_mapping = DreamXWorldARConfig().param_names_mapping
|
||||
reverse_param_names_mapping = DreamXWorldARConfig().reverse_param_names_mapping
|
||||
lora_param_names_mapping = DreamXWorldARConfig().lora_param_names_mapping
|
||||
_no_split_modules = ["CausalWanAttentionBlock"]
|
||||
|
||||
def __init__(self, config: DreamXWorldARConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
model_type = config.model_type
|
||||
patch_size = config.patch_size
|
||||
text_len = config.text_len
|
||||
in_dim = config.in_channels
|
||||
dim = config.hidden_size
|
||||
ffn_dim = config.ffn_dim
|
||||
freq_dim = config.freq_dim
|
||||
text_dim = config.text_dim
|
||||
out_dim = config.out_channels
|
||||
num_heads = config.num_attention_heads
|
||||
num_layers = config.num_layers
|
||||
local_attn_size = config.local_attn_size
|
||||
sink_size = config.sink_size
|
||||
qk_norm = bool(config.qk_norm)
|
||||
cross_attn_norm = config.cross_attn_norm
|
||||
eps = config.eps
|
||||
add_control_adapter = config.add_control_adapter
|
||||
cam_method = config.cam_method
|
||||
attn_compress = config.attn_compress
|
||||
cam_self_attn_layers = config.cam_self_attn_layers
|
||||
|
||||
assert model_type in ['t2v', 'i2v', 'ti2v']
|
||||
self.model_type = model_type
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
self.in_dim = in_dim
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.freq_dim = freq_dim
|
||||
self.text_dim = text_dim
|
||||
self.out_dim = out_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.local_attn_size = local_attn_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
# embeddings — nn.Linear inside nn.Sequential on purpose: the official
|
||||
# checkpoint keys are positional (text_embedding.0/.2, time_embedding.0/.2,
|
||||
# time_projection.1) and must load verbatim (see ffn comment above).
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.text_embedding = nn.Sequential(
|
||||
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(dim, dim))
|
||||
self.time_embedding = nn.Sequential(
|
||||
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
|
||||
# transformer blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
CausalWanAttentionBlock(
|
||||
dim, ffn_dim, num_heads, local_attn_size, sink_size,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
add_control_adapter=add_control_adapter,
|
||||
cam_method=cam_method,
|
||||
attn_compress=attn_compress,
|
||||
layer_idx=layer_idx,
|
||||
cam_self_attn_layers=cam_self_attn_layers)
|
||||
for layer_idx in range(num_layers)
|
||||
])
|
||||
for layer_idx, block in enumerate(self.blocks):
|
||||
block.self_attn.layer_idx = layer_idx
|
||||
block.self_attn.num_layers = self.num_layers
|
||||
|
||||
# head
|
||||
self.head = CausalHead(dim, out_dim, patch_size, eps)
|
||||
|
||||
# RoPE frequencies
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
d = dim // num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
], dim=1)
|
||||
|
||||
self.num_attention_heads = num_heads
|
||||
self.attention_head_dim = dim // num_heads
|
||||
self.hidden_size = dim
|
||||
self.in_channels = in_dim
|
||||
self.out_channels = out_dim
|
||||
self.num_channels_latents = out_dim
|
||||
self.init_weights()
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self, x=None, t=None, context=None, seq_len=None, y=None, y_camera=None,
|
||||
kv_cache=None, crossattn_cache=None, current_start=0,
|
||||
cache_start=0, cache_update_policy="commit_detached",
|
||||
hidden_states=None, encoder_hidden_states=None, timestep=None, **kwargs):
|
||||
"""
|
||||
Causal inference with KV caching.
|
||||
See Algorithm 2 of CausVid (https://arxiv.org/abs/2412.07772).
|
||||
|
||||
Args:
|
||||
x: List of input video tensors [C_in, F, H, W]
|
||||
t: Timestep tensor [B, L]
|
||||
context: List of text embeddings [L, C]
|
||||
seq_len: Maximum sequence length for positional encoding
|
||||
y: Optional conditional video inputs (I2V mode)
|
||||
y_camera: Camera parameters dict {'viewmats': ..., 'K': ...}
|
||||
kv_cache: List of KV cache dicts per transformer block
|
||||
crossattn_cache: List of cross-attention cache dicts
|
||||
current_start: Current position in global token sequence
|
||||
cache_start: Cache start position
|
||||
cache_update_policy: Cache update strategy ('commit_detached' or 'none')
|
||||
|
||||
Returns:
|
||||
Stacked output tensors [B, C_out, F, H/8, W/8]
|
||||
"""
|
||||
if x is None and hidden_states is not None:
|
||||
x = [sample for sample in hidden_states]
|
||||
if t is None and timestep is not None:
|
||||
t = timestep
|
||||
if context is None and encoder_hidden_states is not None:
|
||||
if isinstance(encoder_hidden_states, torch.Tensor):
|
||||
context = [sample for sample in encoder_hidden_states]
|
||||
else:
|
||||
context = encoder_hidden_states
|
||||
if seq_len is None:
|
||||
if torch.is_tensor(t):
|
||||
seq_len = int(t.shape[1]) if t.dim() > 1 else int(t.numel())
|
||||
elif x is not None:
|
||||
sample = x[0]
|
||||
seq_len = (sample.shape[1] // self.patch_size[0]) * (sample.shape[2] // self.patch_size[1]) * (sample.shape[3] // self.patch_size[2])
|
||||
if x is None or t is None or context is None or seq_len is None:
|
||||
raise ValueError("DreamXWorldARTransformer3DModel requires x/t/context/seq_len or FastVideo aliases")
|
||||
|
||||
device = self.patch_embedding.weight.device
|
||||
if self.freqs.is_meta or self.freqs.device != device:
|
||||
d = self.dim // self.num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
], dim=1).to(device)
|
||||
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y, strict=True)]
|
||||
|
||||
# patch embedding
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat(x)
|
||||
|
||||
# time embedding
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t.flatten()).type_as(x))
|
||||
e0 = self.time_projection(e).unflatten(
|
||||
1, (6, self.dim)).unflatten(dim=0, sizes=t.shape)
|
||||
|
||||
# text embedding
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat(
|
||||
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]))
|
||||
|
||||
# camera parameters
|
||||
if y_camera is not None and isinstance(y_camera, dict):
|
||||
cam_viewmats = y_camera['viewmats']
|
||||
cam_K = y_camera['K']
|
||||
else:
|
||||
cam_viewmats = None
|
||||
cam_K = None
|
||||
|
||||
block_kwargs = dict(
|
||||
e=e0, seq_lens=seq_lens, grid_sizes=grid_sizes, freqs=self.freqs,
|
||||
context=context, context_lens=context_lens,
|
||||
cam_viewmats=cam_viewmats, cam_K=cam_K,
|
||||
cache_update_policy=cache_update_policy,
|
||||
)
|
||||
|
||||
cache_update_infos = []
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
block_kwargs.update({
|
||||
"kv_cache": kv_cache[block_index] if kv_cache is not None else None,
|
||||
"crossattn_cache": crossattn_cache[block_index] if crossattn_cache is not None else None,
|
||||
"current_start": current_start,
|
||||
"cache_start": cache_start,
|
||||
})
|
||||
x, block_cache_update_info = block(x, **block_kwargs)
|
||||
if kv_cache is not None:
|
||||
cache_update_infos.append((block_index, block_cache_update_info))
|
||||
|
||||
# Apply deferred cache updates
|
||||
if kv_cache is not None and cache_update_infos and cache_update_policy != "none":
|
||||
self._apply_cache_updates(kv_cache, cache_update_infos)
|
||||
|
||||
# head & unpatchify
|
||||
x = self.head(x, e.unflatten(dim=0, sizes=t.shape).unsqueeze(2))
|
||||
x = self.unpatchify(x, grid_sizes)
|
||||
return torch.stack(x)
|
||||
|
||||
def _apply_cache_updates(self, kv_cache, cache_update_infos):
|
||||
"""Apply deferred cache updates collected from all transformer blocks.
|
||||
|
||||
For Block-Relativistic RoPE, this stores un-roped K values in the cache.
|
||||
RoPE is applied dynamically during attention based on each token's current
|
||||
relative position in the sliding window.
|
||||
"""
|
||||
with torch.no_grad():
|
||||
for block_index, (current_end, local_end_index, update_info) in cache_update_infos:
|
||||
if update_info is not None:
|
||||
cache = kv_cache[block_index]
|
||||
|
||||
if update_info["action"] == "roll_and_insert":
|
||||
sink_tokens = update_info["sink_tokens"]
|
||||
num_rolled_tokens = update_info["num_rolled_tokens"]
|
||||
num_evicted_tokens = update_info["num_evicted_tokens"]
|
||||
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
|
||||
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
|
||||
new_k = update_info["new_k"].detach()
|
||||
new_v = update_info["new_v"].detach()
|
||||
|
||||
cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
|
||||
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
|
||||
cache["k"][:, write_start_index:write_end_index] = new_k
|
||||
cache["v"][:, write_start_index:write_end_index] = new_v
|
||||
|
||||
elif update_info["action"] == "direct_insert":
|
||||
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
|
||||
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
|
||||
new_k = update_info["new_k"].detach()
|
||||
new_v = update_info["new_v"].detach()
|
||||
|
||||
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
|
||||
cache["k"][:, write_start_index:write_end_index] = new_k
|
||||
cache["v"][:, write_start_index:write_end_index] = new_v
|
||||
|
||||
is_recompute = False if update_info is None else update_info.get("is_recompute", False)
|
||||
if not is_recompute:
|
||||
kv_cache[block_index]["global_end_index"].fill_(current_end)
|
||||
kv_cache[block_index]["local_end_index"].fill_(local_end_index)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
"""Reconstruct video tensors from patch embeddings."""
|
||||
c = self.out_dim
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist(), strict=True):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = torch.einsum('fhwpqrc->cfphqwr', u)
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size, strict=True)])
|
||||
out.append(u)
|
||||
return out
|
||||
|
||||
def init_weights(self):
|
||||
"""Initialize model parameters using Xavier initialization."""
|
||||
for m in self.modules():
|
||||
if isinstance(m, (nn.Linear, ReplicatedLinear)):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
|
||||
for m in self.text_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=.02)
|
||||
for m in self.time_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=.02)
|
||||
|
||||
nn.init.zeros_(self.head.head.weight)
|
||||
|
||||
|
||||
EntryClass = DreamXWorldARTransformer3DModel
|
||||
@@ -0,0 +1,578 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.layers.rotary_embedding import apply_rotary_emb, get_1d_rotary_pos_embed
|
||||
|
||||
from fastvideo.attention import DistributedAttention
|
||||
from fastvideo.configs.models import DiTConfig
|
||||
from fastvideo.forward_context import get_forward_context, set_forward_context
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.sd3 import (
|
||||
CombinedTimestepTextProjEmbeddings,
|
||||
SD3AdaLayerNormContinuous,
|
||||
SD3AdaLayerNormZero,
|
||||
SD3FeedForward,
|
||||
SD3TextProjection,
|
||||
SD3TimestepEmbedding,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxTransformer2DModelOutput:
|
||||
sample: torch.Tensor
|
||||
|
||||
|
||||
class FluxPosEmbed(nn.Module):
|
||||
"""1D RoPE axes concatenated per Diffusers `FluxPosEmbed`."""
|
||||
|
||||
def __init__(self, theta: int, axes_dim: list[int]) -> None:
|
||||
super().__init__()
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
|
||||
def forward(self, ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
n_axes = ids.shape[-1]
|
||||
cos_out: list[torch.Tensor] = []
|
||||
sin_out: list[torch.Tensor] = []
|
||||
pos = ids.float()
|
||||
is_mps = ids.device.type == "mps"
|
||||
is_npu = ids.device.type == "npu"
|
||||
freqs_dtype = torch.float32 if (is_mps or is_npu) else torch.float64
|
||||
for i in range(n_axes):
|
||||
cos, sin = get_1d_rotary_pos_embed(
|
||||
self.axes_dim[i],
|
||||
pos[:, i],
|
||||
theta=self.theta,
|
||||
use_real=True,
|
||||
freqs_dtype=freqs_dtype,
|
||||
)
|
||||
cos_out.append(cos)
|
||||
sin_out.append(sin)
|
||||
freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device)
|
||||
freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
class FluxCombinedTimestepGuidanceTextProjEmbeddings(nn.Module):
|
||||
def __init__(self, embedding_dim: int, pooled_projection_dim: int) -> None:
|
||||
super().__init__()
|
||||
self.time_proj = Timesteps(
|
||||
num_channels=256,
|
||||
flip_sin_to_cos=True,
|
||||
downscale_freq_shift=0,
|
||||
)
|
||||
self.timestep_embedder = SD3TimestepEmbedding(
|
||||
in_channels=256,
|
||||
time_embed_dim=embedding_dim,
|
||||
act_fn="silu",
|
||||
)
|
||||
self.guidance_embedder = SD3TimestepEmbedding(
|
||||
in_channels=256,
|
||||
time_embed_dim=embedding_dim,
|
||||
act_fn="silu",
|
||||
)
|
||||
self.text_embedder = SD3TextProjection(
|
||||
pooled_projection_dim,
|
||||
embedding_dim,
|
||||
act_fn="silu",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
guidance: torch.Tensor,
|
||||
pooled_projection: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype))
|
||||
guidance_proj = self.time_proj(guidance)
|
||||
guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=pooled_projection.dtype))
|
||||
time_guidance_emb = timesteps_emb + guidance_emb
|
||||
pooled_projections = self.text_embedder(pooled_projection)
|
||||
return time_guidance_emb + pooled_projections
|
||||
|
||||
|
||||
class FluxAdaLayerNormZeroSingle(nn.Module):
|
||||
def __init__(self, embedding_dim: int, bias: bool = True) -> None:
|
||||
super().__init__()
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = ReplicatedLinear(embedding_dim, 3 * embedding_dim, bias=bias)
|
||||
self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
emb: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
emb, _ = self.linear(self.silu(emb))
|
||||
shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1)
|
||||
x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None]
|
||||
return x, gate_msa
|
||||
|
||||
|
||||
class FluxJointAttention(nn.Module):
|
||||
"""Joint attention: text tokens precede image tokens (Diffusers order)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.heads = num_attention_heads
|
||||
self.head_dim = attention_head_dim
|
||||
self.inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_added_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_added_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
|
||||
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.add_q_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.add_k_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.add_v_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
|
||||
self.to_out = nn.ModuleList(
|
||||
[
|
||||
ReplicatedLinear(self.inner_dim, dim, bias=True),
|
||||
nn.Dropout(0.0),
|
||||
]
|
||||
)
|
||||
self.to_add_out = ReplicatedLinear(self.inner_dim, dim, bias=True)
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=attention_head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
batch_size = hidden_states.shape[0]
|
||||
text_seq_len = encoder_hidden_states.shape[1]
|
||||
img_seq_len = hidden_states.shape[1]
|
||||
|
||||
q, _ = self.to_q(hidden_states)
|
||||
k, _ = self.to_k(hidden_states)
|
||||
v, _ = self.to_v(hidden_states)
|
||||
q = q.view(batch_size, img_seq_len, self.heads, self.head_dim)
|
||||
k = k.view(batch_size, img_seq_len, self.heads, self.head_dim)
|
||||
v = v.view(batch_size, img_seq_len, self.heads, self.head_dim)
|
||||
q = self.norm_q(q)
|
||||
k = self.norm_k(k)
|
||||
|
||||
enc_q, _ = self.add_q_proj(encoder_hidden_states)
|
||||
enc_k, _ = self.add_k_proj(encoder_hidden_states)
|
||||
enc_v, _ = self.add_v_proj(encoder_hidden_states)
|
||||
enc_q = enc_q.view(batch_size, text_seq_len, self.heads, self.head_dim)
|
||||
enc_k = enc_k.view(batch_size, text_seq_len, self.heads, self.head_dim)
|
||||
enc_v = enc_v.view(batch_size, text_seq_len, self.heads, self.head_dim)
|
||||
enc_q = self.norm_added_q(enc_q)
|
||||
enc_k = self.norm_added_k(enc_k)
|
||||
|
||||
q = torch.cat([enc_q, q], dim=1)
|
||||
k = torch.cat([enc_k, k], dim=1)
|
||||
v = torch.cat([enc_v, v], dim=1)
|
||||
|
||||
q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1)
|
||||
k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1)
|
||||
|
||||
joint_out, _ = self.attn(q, k, v)
|
||||
joint_out = joint_out.reshape(batch_size, text_seq_len + img_seq_len, self.inner_dim)
|
||||
|
||||
enc_out = joint_out[:, :text_seq_len]
|
||||
img_out = joint_out[:, text_seq_len:]
|
||||
|
||||
img_out, _ = self.to_out[0](img_out)
|
||||
img_out = self.to_out[1](img_out)
|
||||
enc_out, _ = self.to_add_out(enc_out)
|
||||
return img_out, enc_out
|
||||
|
||||
|
||||
class FluxSingleStreamAttention(nn.Module):
|
||||
"""Self-attention on concatenated text+image sequence (single blocks)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.heads = num_attention_heads
|
||||
self.head_dim = attention_head_dim
|
||||
self.inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=attention_head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
batch_size, seq_len, _ = hidden_states.shape
|
||||
q, _ = self.to_q(hidden_states)
|
||||
k, _ = self.to_k(hidden_states)
|
||||
v, _ = self.to_v(hidden_states)
|
||||
q = q.view(batch_size, seq_len, self.heads, self.head_dim)
|
||||
k = k.view(batch_size, seq_len, self.heads, self.head_dim)
|
||||
v = v.view(batch_size, seq_len, self.heads, self.head_dim)
|
||||
q = self.norm_q(q)
|
||||
k = self.norm_k(k)
|
||||
q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1)
|
||||
k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1)
|
||||
out, _ = self.attn(q, k, v)
|
||||
return out.reshape(batch_size, seq_len, self.inner_dim)
|
||||
|
||||
|
||||
class FluxTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.norm1 = SD3AdaLayerNormZero(dim)
|
||||
self.norm1_context = SD3AdaLayerNormZero(dim)
|
||||
self.attn = FluxJointAttention(
|
||||
dim=dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
||||
self.ff = SD3FeedForward(
|
||||
dim=dim,
|
||||
dim_out=dim,
|
||||
activation_fn="gelu-approximate",
|
||||
)
|
||||
self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
||||
self.ff_context = SD3FeedForward(
|
||||
dim=dim,
|
||||
dim_out=dim,
|
||||
activation_fn="gelu-approximate",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
joint_attention_kwargs: dict[str, Any] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del joint_attention_kwargs
|
||||
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb)
|
||||
(norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp) = self.norm1_context(
|
||||
encoder_hidden_states, emb=temb
|
||||
)
|
||||
|
||||
attn_output, context_attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
|
||||
attn_output = gate_msa.unsqueeze(1) * attn_output
|
||||
hidden_states = hidden_states + attn_output
|
||||
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output
|
||||
|
||||
context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output
|
||||
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
||||
|
||||
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
||||
norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None]
|
||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||
encoder_hidden_states = encoder_hidden_states + (c_gate_mlp.unsqueeze(1) * context_ff_output)
|
||||
if encoder_hidden_states.dtype == torch.float16:
|
||||
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
||||
|
||||
return encoder_hidden_states, hidden_states
|
||||
|
||||
|
||||
class FluxSingleTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.norm = FluxAdaLayerNormZeroSingle(dim)
|
||||
self.proj_mlp = ReplicatedLinear(dim, mlp_hidden_dim, bias=True)
|
||||
self.act_mlp = nn.GELU(approximate="tanh")
|
||||
self.proj_out = ReplicatedLinear(dim + mlp_hidden_dim, dim, bias=True)
|
||||
self.attn = FluxSingleStreamAttention(
|
||||
dim=dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
joint_attention_kwargs: dict[str, Any] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del joint_attention_kwargs
|
||||
text_seq_len = encoder_hidden_states.shape[1]
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
residual = hidden_states
|
||||
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
||||
mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)[0])
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
||||
gate = gate.unsqueeze(1)
|
||||
hidden_states = gate * self.proj_out(hidden_states)[0]
|
||||
hidden_states = residual + hidden_states
|
||||
if hidden_states.dtype == torch.float16:
|
||||
hidden_states = hidden_states.clip(-65504, 65504)
|
||||
encoder_hidden_states = hidden_states[:, :text_seq_len]
|
||||
hidden_states = hidden_states[:, text_seq_len:]
|
||||
return encoder_hidden_states, hidden_states
|
||||
|
||||
|
||||
class FluxTransformer2DModel(BaseDiT):
|
||||
"""FastVideo FLUX transformer; load Diffusers FLUX safetensors 1:1."""
|
||||
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: (n.startswith("transformer_blocks.") or n.startswith("single_transformer_blocks."))
|
||||
and n.split(".")[-1].isdigit(),
|
||||
]
|
||||
_compile_conditions = _fsdp_shard_conditions
|
||||
# HF weight names already match this module layout (cf. SGLang regex maps).
|
||||
param_names_mapping: dict[str, Any] = {}
|
||||
reverse_param_names_mapping: dict[str, Any] = {}
|
||||
lora_param_names_mapping: dict[str, Any] = {}
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs) -> None:
|
||||
del kwargs
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
self.fastvideo_config = config
|
||||
self.hf_config = hf_config
|
||||
arch = config.arch_config
|
||||
|
||||
out_ch = arch.out_channels
|
||||
self.out_channels = out_ch if out_ch is not None else arch.in_channels
|
||||
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
|
||||
self.hidden_size = self.inner_dim
|
||||
self.num_attention_heads = arch.num_attention_heads
|
||||
self.num_channels_latents = arch.in_channels
|
||||
|
||||
axes_list = list(arch.axes_dims_rope)
|
||||
self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_list)
|
||||
if arch.guidance_embeds:
|
||||
self.time_text_embed = FluxCombinedTimestepGuidanceTextProjEmbeddings(
|
||||
embedding_dim=self.inner_dim,
|
||||
pooled_projection_dim=arch.pooled_projection_dim,
|
||||
)
|
||||
else:
|
||||
self.time_text_embed = CombinedTimestepTextProjEmbeddings(
|
||||
embedding_dim=self.inner_dim,
|
||||
pooled_projection_dim=arch.pooled_projection_dim,
|
||||
)
|
||||
self.context_embedder = ReplicatedLinear(arch.joint_attention_dim, self.inner_dim)
|
||||
self.x_embedder = ReplicatedLinear(arch.in_channels, self.inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
FluxTransformerBlock(
|
||||
dim=self.inner_dim,
|
||||
num_attention_heads=arch.num_attention_heads,
|
||||
attention_head_dim=arch.attention_head_dim,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
)
|
||||
for _ in range(arch.num_layers)
|
||||
]
|
||||
)
|
||||
self.single_transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
FluxSingleTransformerBlock(
|
||||
dim=self.inner_dim,
|
||||
num_attention_heads=arch.num_attention_heads,
|
||||
attention_head_dim=arch.attention_head_dim,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
)
|
||||
for _ in range(arch.num_single_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm_out = SD3AdaLayerNormContinuous(
|
||||
self.inner_dim,
|
||||
self.inner_dim,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
bias=True,
|
||||
norm_type="layer_norm",
|
||||
)
|
||||
self.proj_out = ReplicatedLinear(
|
||||
self.inner_dim,
|
||||
arch.patch_size * arch.patch_size * self.out_channels,
|
||||
bias=True,
|
||||
)
|
||||
self.gradient_checkpointing = False
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | None = None,
|
||||
pooled_projections: torch.Tensor | None = None,
|
||||
timestep: torch.LongTensor | torch.Tensor | None = None,
|
||||
img_ids: torch.Tensor | None = None,
|
||||
txt_ids: torch.Tensor | None = None,
|
||||
guidance: torch.Tensor | None = None,
|
||||
joint_attention_kwargs: dict[str, Any] | None = None,
|
||||
return_dict: bool = True,
|
||||
controlnet_block_samples: Any | None = None,
|
||||
controlnet_single_block_samples: Any | None = None,
|
||||
controlnet_blocks_repeat: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> FluxTransformer2DModelOutput | tuple[torch.Tensor, ...]:
|
||||
del kwargs
|
||||
if encoder_hidden_states is None:
|
||||
raise ValueError("encoder_hidden_states must be provided")
|
||||
if pooled_projections is None:
|
||||
raise ValueError("pooled_projections must be provided")
|
||||
if timestep is None:
|
||||
raise ValueError("timestep must be provided")
|
||||
if img_ids is None or txt_ids is None:
|
||||
raise ValueError("img_ids and txt_ids must be provided")
|
||||
|
||||
arch = self.fastvideo_config.arch_config
|
||||
if arch.guidance_embeds and guidance is None:
|
||||
raise ValueError("guidance must be provided when guidance_embeds=True")
|
||||
|
||||
if timestep.dim() == 0:
|
||||
timestep = timestep[None]
|
||||
if timestep.dim() > 1:
|
||||
timestep = timestep.reshape(-1)
|
||||
if timestep.shape[0] == 1 and hidden_states.shape[0] > 1:
|
||||
timestep = timestep.expand(hidden_states.shape[0])
|
||||
|
||||
try:
|
||||
get_forward_context()
|
||||
forward_context = nullcontext()
|
||||
except AssertionError:
|
||||
if timestep.numel() == 0:
|
||||
ts0 = 0
|
||||
elif torch.is_floating_point(timestep):
|
||||
ts0 = int(round(timestep[0].item() * 1000))
|
||||
else:
|
||||
ts0 = int(timestep[0].item())
|
||||
forward_context = set_forward_context(current_timestep=ts0, attn_metadata=None)
|
||||
|
||||
with forward_context:
|
||||
hidden_states, _ = self.x_embedder(hidden_states)
|
||||
|
||||
ts = timestep.to(hidden_states.dtype) * 1000
|
||||
g = None if guidance is None else guidance.to(hidden_states.dtype) * 1000
|
||||
|
||||
if arch.guidance_embeds:
|
||||
assert g is not None
|
||||
temb = self.time_text_embed(ts, g, pooled_projections)
|
||||
else:
|
||||
temb = self.time_text_embed(timestep=ts, pooled_projection=pooled_projections)
|
||||
|
||||
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
|
||||
|
||||
if txt_ids.ndim == 3:
|
||||
txt_ids = txt_ids[0]
|
||||
if img_ids.ndim == 3:
|
||||
img_ids = img_ids[0]
|
||||
|
||||
ids = torch.cat((txt_ids, img_ids), dim=0)
|
||||
image_rotary_emb = self.pos_embed(ids)
|
||||
|
||||
jkwargs = joint_attention_kwargs or {}
|
||||
|
||||
for idx, block in enumerate(self.transformer_blocks):
|
||||
encoder_hidden_states, hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
joint_attention_kwargs=jkwargs,
|
||||
)
|
||||
if controlnet_block_samples:
|
||||
interval = len(self.transformer_blocks) / len(controlnet_block_samples)
|
||||
interval = int(math.ceil(interval))
|
||||
if controlnet_blocks_repeat:
|
||||
hidden_states = hidden_states + controlnet_block_samples[idx % len(controlnet_block_samples)]
|
||||
else:
|
||||
hidden_states = hidden_states + controlnet_block_samples[idx // interval]
|
||||
|
||||
for idx, block in enumerate(self.single_transformer_blocks):
|
||||
encoder_hidden_states, hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
joint_attention_kwargs=jkwargs,
|
||||
)
|
||||
if controlnet_single_block_samples:
|
||||
interval = len(self.single_transformer_blocks) / len(controlnet_single_block_samples)
|
||||
interval = int(math.ceil(interval))
|
||||
hidden_states = hidden_states + controlnet_single_block_samples[idx // interval]
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
output, _ = self.proj_out(hidden_states)
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
return FluxTransformer2DModelOutput(sample=output)
|
||||
|
||||
|
||||
EntryClass = FluxTransformer2DModel
|
||||
@@ -699,8 +699,8 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
if mask is None:
|
||||
attn_output = self.attn(q, k, v)
|
||||
else:
|
||||
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
|
||||
attn_metadata = FlashAttnMetadataBuilder().build(
|
||||
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
|
||||
attn_metadata = SDPAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
attn_mask=mask,
|
||||
)
|
||||
|
||||
@@ -173,8 +173,11 @@ class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
|
||||
img_v_prope = img_v_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
|
||||
# end hyworld
|
||||
|
||||
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
|
||||
attn_metadata = FlashAttnMetadataBuilder().build(
|
||||
# The metadata only carries the text padding mask; the executing
|
||||
# attention kernel is still whatever the layer's selector picked
|
||||
# (flash-attn when installed), not necessarily torch SDPA.
|
||||
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
|
||||
attn_metadata = SDPAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
attn_mask=encoder_attention_mask,
|
||||
)
|
||||
@@ -192,14 +195,9 @@ class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
|
||||
)
|
||||
|
||||
# begin hyworld
|
||||
# attention with prope
|
||||
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
|
||||
attn_metadata_prope = FlashAttnMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
attn_mask=encoder_attention_mask,
|
||||
)
|
||||
# attention with prope (same text mask, so reuse the metadata)
|
||||
# NOTE: Do NOT pass freqs_cis to prope attention - HY-WorldPlay does not apply RoPE to prope
|
||||
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata_prope):
|
||||
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
|
||||
img_attn_prope, _ = self.attn(
|
||||
img_q_prope,
|
||||
img_k_prope,
|
||||
|
||||
@@ -10,13 +10,21 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.attention.backends.nabla import CAN_USE_FLEX_ATTN, flex_attention, nablaT_v2
|
||||
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
|
||||
from fastvideo.layers.layernorm import LayerNormScaleShift
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
if not CAN_USE_FLEX_ATTN:
|
||||
logger.warning("torch.nn.attention.flex_attention is unavailable in this PyTorch build; "
|
||||
"Kandinsky5 NABLA sparse attention (Pro checkpoints) cannot be used.")
|
||||
|
||||
FRACTAL_PIXEL_SIZE = 8
|
||||
_ARCH_CONFIG_DEFAULTS = Kandinsky5VideoConfig().arch_config
|
||||
|
||||
@@ -263,10 +271,9 @@ class Kandinsky5Modulation(nn.Module):
|
||||
|
||||
|
||||
def _apply_rotary(x: torch.Tensor, rope: torch.Tensor) -> torch.Tensor:
|
||||
orig_dtype = x.dtype
|
||||
x_ = x.reshape(*x.shape[:-1], -1, 1, 2).to(torch.float32)
|
||||
x_out = (rope * x_).sum(dim=-1)
|
||||
return x_out.reshape(*x.shape).to(orig_dtype)
|
||||
return x_out.reshape(*x.shape).to(x.dtype)
|
||||
|
||||
|
||||
class Kandinsky5Attention(nn.Module):
|
||||
@@ -277,6 +284,7 @@ class Kandinsky5Attention(nn.Module):
|
||||
head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None,
|
||||
prefix: str = "",
|
||||
use_nabla: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
assert num_channels % head_dim == 0
|
||||
@@ -306,6 +314,17 @@ class Kandinsky5Attention(nn.Module):
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
# NABLA checkpoints get a second attention layer whose backend defaults
|
||||
# to NABLA_ATTN; FASTVIDEO_ATTENTION_BACKEND still overrides it.
|
||||
self.nabla_attention = None
|
||||
if use_nabla:
|
||||
self.nabla_attention = LocalAttention(
|
||||
num_heads=self.num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
default_backend=AttentionBackendEnum.NABLA_ATTN,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -339,27 +358,54 @@ class Kandinsky5Attention(nn.Module):
|
||||
key = _apply_rotary(key, rotary_emb).type_as(key)
|
||||
|
||||
if sparse_params is not None:
|
||||
raise NotImplementedError(
|
||||
"Sparse attention is not yet supported for Kandinsky5 in FastVideo."
|
||||
)
|
||||
if self.nabla_attention is None:
|
||||
raise RuntimeError("sparse_params passed to an attention layer built without use_nabla; "
|
||||
"this checkpoint/config combination is inconsistent.")
|
||||
try:
|
||||
# Backend impl reads sta_mask/P from the forward-context
|
||||
# attention metadata built by the denoising stage.
|
||||
hidden_states = self.nabla_attention(query, key, value)
|
||||
except AssertionError as exc:
|
||||
# Standalone parity tests call the model without a pipeline
|
||||
# forward context; run the NABLA kernel directly.
|
||||
if "Forward context is not set" not in str(exc):
|
||||
raise
|
||||
attn_mask = nablaT_v2(query, key, sparse_params["sta_mask"], thr=sparse_params["P"])
|
||||
hidden_states = flex_attention(
|
||||
query=query.transpose(1, 2),
|
||||
key=key.transpose(1, 2),
|
||||
value=value.transpose(1, 2),
|
||||
block_mask=attn_mask,
|
||||
).transpose(1, 2)
|
||||
else:
|
||||
try:
|
||||
hidden_states = self.local_attention(query, key, value)
|
||||
|
||||
try:
|
||||
hidden_states = self.local_attention(query, key, value)
|
||||
except AssertionError as exc:
|
||||
# LocalAttention requires pipeline forward context. Standalone
|
||||
# parity tests call the model directly, so fallback to Torch SDPA.
|
||||
if "Forward context is not set" not in str(exc):
|
||||
raise
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
hidden_states = F.scaled_dot_product_attention(query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2)
|
||||
hidden_states = hidden_states.flatten(2)
|
||||
except AssertionError as exc:
|
||||
# LocalAttention requires pipeline forward context. Standalone
|
||||
# parity tests call the model directly, so fallback to Torch SDPA.
|
||||
if "Forward context is not set" not in str(exc):
|
||||
raise
|
||||
|
||||
query_shape = query.shape[:-2]
|
||||
key_shape = key.shape[:-2]
|
||||
query = query.reshape(query_shape[0], -1, self.num_heads,
|
||||
query.shape[-1]).transpose(1, 2)
|
||||
key = key.reshape(key_shape[0], -1, self.num_heads,
|
||||
key.shape[-1]).transpose(1, 2)
|
||||
value = value.reshape(key_shape[0], -1, self.num_heads,
|
||||
value.shape[-1]).transpose(1, 2)
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
is_causal=False,
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(
|
||||
*query_shape, self.num_heads, -1)
|
||||
|
||||
hidden_states = hidden_states.flatten(-2, -1)
|
||||
|
||||
hidden_states, _ = self.out_layer(hidden_states)
|
||||
return hidden_states
|
||||
@@ -476,7 +522,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
|
||||
head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
prefix: str = ""):
|
||||
prefix: str = "",
|
||||
use_nabla: bool = False):
|
||||
super().__init__()
|
||||
self.visual_modulation = Kandinsky5Modulation(time_dim, model_dim, 9)
|
||||
|
||||
@@ -491,7 +538,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
|
||||
model_dim,
|
||||
head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.self_attention")
|
||||
prefix=f"{prefix}.self_attention",
|
||||
use_nabla=use_nabla)
|
||||
|
||||
self.cross_attention_norm = LayerNormScaleShift(
|
||||
model_dim,
|
||||
@@ -624,7 +672,8 @@ class Kandinsky5Transformer3DModel(BaseDiT):
|
||||
arch.ff_dim,
|
||||
head_dim,
|
||||
self._supported_attention_backends,
|
||||
prefix=f"{config.prefix}.visual_transformer_blocks.{i}")
|
||||
prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
|
||||
use_nabla=arch.attention_type == "nabla")
|
||||
for i in range(arch.num_visual_blocks)
|
||||
])
|
||||
|
||||
@@ -694,6 +743,7 @@ class Kandinsky5Transformer3DModel(BaseDiT):
|
||||
scale_factor)
|
||||
to_fractal = sparse_params[
|
||||
"to_fractal"] if sparse_params is not None else False
|
||||
|
||||
visual_embed, visual_rope = fractal_flatten(visual_embed, visual_rope,
|
||||
visual_shape,
|
||||
block_mask=to_fractal)
|
||||
@@ -724,6 +774,7 @@ class Kandinsky5Transformer3DModel(BaseDiT):
|
||||
|
||||
if return_dict:
|
||||
return Kandinsky5TransformerOutput(sample=x)
|
||||
|
||||
return x
|
||||
|
||||
def materialize_non_persistent_buffers(self, device: torch.device,
|
||||
|
||||
@@ -16,6 +16,8 @@ from fastvideo.layers.rotary_embedding import (
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
from .utils import retain_kv_with_sink
|
||||
|
||||
|
||||
DISABLE_COMPILE = False
|
||||
flex_attention = torch.compile(
|
||||
@@ -100,6 +102,7 @@ def _update_kv_cache_and_attend(
|
||||
start_frame: int,
|
||||
num_frame_per_block: int,
|
||||
local_attn_size: int,
|
||||
sink_size: int,
|
||||
*,
|
||||
use_k_for_num_tokens: bool = False,
|
||||
store_first_only: bool = False,
|
||||
@@ -133,7 +136,6 @@ def _update_kv_cache_and_attend(
|
||||
k.shape[1] if use_k_for_num_tokens else q.shape[1]
|
||||
) == num_frame_per_block
|
||||
|
||||
sink_size = 0
|
||||
max_attention_size = local_attn_size
|
||||
sink_tokens = sink_size * 1
|
||||
kv_cache_size = kv_cache["k"].shape[1]
|
||||
@@ -206,16 +208,22 @@ def _update_kv_cache_and_attend(
|
||||
|
||||
# Store new k, v in cache
|
||||
if store_first_only:
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = k[:1]
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v[:1]
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = k[
|
||||
: kv_cache["k"].shape[0]
|
||||
]
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v[
|
||||
: kv_cache["v"].shape[0]
|
||||
]
|
||||
else:
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = k
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
|
||||
# Retrieve from cache and perform attention
|
||||
cache_start = max(0, local_end_index - max_attention_size)
|
||||
cached_k = kv_cache["k"][:, cache_start:local_end_index]
|
||||
cached_v = kv_cache["v"][:, cache_start:local_end_index]
|
||||
cached_k, cached_v = retain_kv_with_sink(
|
||||
kv_cache["k"][:, :local_end_index],
|
||||
kv_cache["v"][:, :local_end_index],
|
||||
min(local_end_index, max_attention_size),
|
||||
sink_tokens,
|
||||
)
|
||||
|
||||
if repeat_factor is not None:
|
||||
cached_k = cached_k.repeat(repeat_factor, 1, 1, 1)
|
||||
@@ -261,6 +269,7 @@ class ActionModule(nn.Module):
|
||||
enable_mouse=True,
|
||||
enable_keyboard=True,
|
||||
local_attn_size=6,
|
||||
sink_size=0,
|
||||
blocks: list | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -274,6 +283,7 @@ class ActionModule(nn.Module):
|
||||
)
|
||||
blocks = blocks if blocks is not None else []
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.enable_mouse = enable_mouse
|
||||
self.enable_keyboard = enable_keyboard
|
||||
|
||||
@@ -565,6 +575,7 @@ class ActionModule(nn.Module):
|
||||
start_frame,
|
||||
num_frame_per_block,
|
||||
self.local_attn_size,
|
||||
self.sink_size,
|
||||
use_k_for_num_tokens=False,
|
||||
)
|
||||
else:
|
||||
@@ -689,6 +700,7 @@ class ActionModule(nn.Module):
|
||||
start_frame,
|
||||
num_frame_per_block,
|
||||
self.local_attn_size,
|
||||
self.sink_size,
|
||||
use_k_for_num_tokens=True,
|
||||
store_first_only=True,
|
||||
repeat_factor=S,
|
||||
@@ -719,6 +731,7 @@ class ActionModule(nn.Module):
|
||||
start_frame,
|
||||
num_frame_per_block,
|
||||
self.local_attn_size,
|
||||
self.sink_size,
|
||||
use_k_for_num_tokens=True,
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -46,6 +46,7 @@ from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
from .action_module import ActionModule
|
||||
from .model import MatrixGame2CrossAttention
|
||||
from .utils import retain_kv_with_sink
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -345,9 +346,12 @@ class CausalMatrixGame2SelfAttention(nn.Module):
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = stored_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
|
||||
kv_start = max(0, local_end_index - max_attention_size)
|
||||
k_for_attn = kv_cache["k"][:, kv_start:local_end_index]
|
||||
v_for_attn = kv_cache["v"][:, kv_start:local_end_index]
|
||||
k_for_attn, v_for_attn = retain_kv_with_sink(
|
||||
kv_cache["k"][:, :local_end_index],
|
||||
kv_cache["v"][:, :local_end_index],
|
||||
min(local_end_index, max_attention_size),
|
||||
sink_tokens,
|
||||
)
|
||||
|
||||
if relativistic:
|
||||
window_len, query_lo, _ = relativistic_window_offsets(
|
||||
@@ -418,6 +422,7 @@ class CausalMatrixGame2TransformerBlock(nn.Module):
|
||||
self.hidden_dim = dim
|
||||
self.num_attention_heads = num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
dim_head = dim // num_heads
|
||||
if qk_norm == "rms_norm":
|
||||
self.norm_q = RMSNorm(dim_head, eps=eps)
|
||||
@@ -474,6 +479,7 @@ class CausalMatrixGame2TransformerBlock(nn.Module):
|
||||
),
|
||||
patch_size=action_config["patch_size"],
|
||||
local_attn_size=local_attn_size,
|
||||
sink_size=sink_size,
|
||||
qk_norm=action_config["qk_norm"],
|
||||
qkv_bias=action_config["qkv_bias"],
|
||||
vae_time_compression_ratio=action_config[
|
||||
@@ -537,9 +543,9 @@ class CausalMatrixGame2TransformerBlock(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))
|
||||
@@ -647,6 +653,8 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
|
||||
# Long tuning controls the causal window from YAML; consume it here
|
||||
# like Wan so train wrappers do not need runtime patching.
|
||||
arch_cfg = getattr(config, "arch_config", None)
|
||||
self.local_attn_size = (
|
||||
getattr(
|
||||
@@ -662,6 +670,12 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
if arch_cfg
|
||||
else getattr(config, "sink_size", 0)
|
||||
)
|
||||
if self.sink_size < 0:
|
||||
raise ValueError("sink_size must be non-negative")
|
||||
if self.local_attn_size != -1 and self.sink_size >= self.local_attn_size:
|
||||
raise ValueError(
|
||||
"sink_size must be smaller than local_attn_size for "
|
||||
"MatrixGame2 causal attention")
|
||||
self.rope_cache_policy = (
|
||||
getattr(arch_cfg, "rope_cache_policy",
|
||||
getattr(config, "rope_cache_policy", "absolute"))
|
||||
@@ -749,6 +763,52 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def set_causal_attention_window(
|
||||
self,
|
||||
*,
|
||||
local_attn_size: int | None = None,
|
||||
sink_size: int | None = None,
|
||||
) -> None:
|
||||
local_attn_size = (
|
||||
self.local_attn_size if local_attn_size is None else int(local_attn_size)
|
||||
)
|
||||
sink_size = self.sink_size if sink_size is None else int(sink_size)
|
||||
if sink_size < 0:
|
||||
raise ValueError("sink_size must be non-negative")
|
||||
if local_attn_size != -1 and sink_size >= local_attn_size:
|
||||
raise ValueError(
|
||||
"sink_size must be smaller than local_attn_size for "
|
||||
"MatrixGame2 causal attention")
|
||||
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.block_mask = None
|
||||
self.block_mask_keyboard = None
|
||||
self.block_mask_mouse = None
|
||||
|
||||
for block in self.blocks:
|
||||
if hasattr(block, "local_attn_size"):
|
||||
block.local_attn_size = local_attn_size
|
||||
if hasattr(block, "sink_size"):
|
||||
block.sink_size = sink_size
|
||||
|
||||
attn1 = getattr(block, "attn1", None)
|
||||
if attn1 is not None:
|
||||
if hasattr(attn1, "local_attn_size"):
|
||||
attn1.local_attn_size = local_attn_size
|
||||
if hasattr(attn1, "sink_size"):
|
||||
attn1.sink_size = sink_size
|
||||
|
||||
action_model = getattr(block, "action_model", None)
|
||||
if action_model is not None:
|
||||
if hasattr(action_model, "local_attn_size"):
|
||||
action_model.local_attn_size = local_attn_size
|
||||
if hasattr(action_model, "sink_size"):
|
||||
action_model.sink_size = sink_size
|
||||
|
||||
def set_sink_size(self, sink_size: int) -> None:
|
||||
self.set_causal_attention_window(sink_size=sink_size)
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
device: torch.device | str,
|
||||
@@ -756,6 +816,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen: int = 880,
|
||||
num_frame_per_block: int = 1,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
) -> BlockMask:
|
||||
total_length = num_frames * frame_seqlen
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
@@ -782,7 +843,10 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
else:
|
||||
return (
|
||||
(kv_idx < ends[q_idx])
|
||||
& (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))
|
||||
& (
|
||||
(kv_idx < sink_size * frame_seqlen)
|
||||
| (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))
|
||||
)
|
||||
) | (q_idx == kv_idx)
|
||||
|
||||
block_mask = create_block_mask(
|
||||
@@ -796,8 +860,9 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(
|
||||
f" cache a block wise causal mask with block size of {num_frame_per_block} frames"
|
||||
logger.info(
|
||||
"cache a block wise causal mask with block size of %s frames",
|
||||
num_frame_per_block,
|
||||
)
|
||||
|
||||
return block_mask
|
||||
@@ -809,6 +874,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen: int = 880,
|
||||
num_frame_per_block: int = 1,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
) -> BlockMask:
|
||||
total_length2 = num_frames * frame_seqlen
|
||||
padded_length2 = math.ceil(total_length2 / 32) * 32 - total_length2
|
||||
@@ -834,7 +900,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
else:
|
||||
return (
|
||||
(kv_idx < ends2[q_idx])
|
||||
& (kv_idx >= (ends2[q_idx] - local_attn_size))
|
||||
& ((kv_idx < sink_size) | (kv_idx >= (ends2[q_idx] - local_attn_size)))
|
||||
) | (q_idx == kv_idx)
|
||||
|
||||
block_mask2 = create_block_mask(
|
||||
@@ -848,8 +914,9 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(
|
||||
f" cache a block wise causal mask for keyboard with block size of {num_frame_per_block} frames"
|
||||
logger.info(
|
||||
"cache a block wise causal mask for keyboard with block size of %s frames",
|
||||
num_frame_per_block,
|
||||
)
|
||||
|
||||
return block_mask2
|
||||
@@ -861,6 +928,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen: int = 1,
|
||||
num_frame_per_block: int = 1,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
) -> BlockMask:
|
||||
total_length2 = num_frames * frame_seqlen
|
||||
padded_length2 = math.ceil(total_length2 / 32) * 32 - total_length2
|
||||
@@ -886,7 +954,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
else:
|
||||
return (
|
||||
(kv_idx < ends2[q_idx])
|
||||
& (kv_idx >= (ends2[q_idx] - local_attn_size))
|
||||
& ((kv_idx < sink_size) | (kv_idx >= (ends2[q_idx] - local_attn_size)))
|
||||
) | (q_idx == kv_idx)
|
||||
|
||||
block_mask2 = create_block_mask(
|
||||
@@ -900,8 +968,9 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(
|
||||
f" cache a block wise causal mask for action with block size of {num_frame_per_block} frames"
|
||||
logger.info(
|
||||
"cache a block wise causal mask for action with block size of %s frames",
|
||||
num_frame_per_block,
|
||||
)
|
||||
|
||||
return block_mask2
|
||||
@@ -1063,6 +1132,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
if self.use_rope_keyboard:
|
||||
block_mask_keyboard = self._prepare_blockwise_causal_attn_mask_action(
|
||||
@@ -1071,6 +1141,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=1,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
else:
|
||||
block_mask_keyboard = self._prepare_blockwise_causal_attn_mask_keyboard(
|
||||
@@ -1079,6 +1150,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
block_mask_mouse = self._prepare_blockwise_causal_attn_mask_action(
|
||||
device=hidden_states.device,
|
||||
@@ -1086,6 +1158,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=1,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
if kv_cache is None:
|
||||
kv_cache = [None] * len(self.blocks)
|
||||
@@ -1239,6 +1312,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
if self.block_mask_keyboard is None:
|
||||
if self.use_rope_keyboard:
|
||||
@@ -1248,6 +1322,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=1,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
else:
|
||||
self.block_mask_keyboard = self._prepare_blockwise_causal_attn_mask_keyboard(
|
||||
@@ -1256,6 +1331,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
if self.block_mask_mouse is None:
|
||||
self.block_mask_mouse = self._prepare_blockwise_causal_attn_mask_action(
|
||||
@@ -1264,6 +1340,7 @@ class CausalMatrixGame2WanModel(BaseDiT):
|
||||
frame_seqlen=1,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
sink_size=self.sink_size,
|
||||
)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
|
||||
@@ -452,7 +452,7 @@ class MatrixGame2WanModel(BaseDiT):
|
||||
and len(encoder_hidden_states_image) > 0
|
||||
):
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
else:
|
||||
elif not isinstance(encoder_hidden_states_image, torch.Tensor):
|
||||
encoder_hidden_states_image = None
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = (
|
||||
|
||||
@@ -4,12 +4,18 @@ import asyncio
|
||||
import os
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.distributed.parallel_state import get_local_torch_device
|
||||
from fastvideo.utils import logger
|
||||
|
||||
try:
|
||||
import cv2
|
||||
except ImportError:
|
||||
cv2 = None
|
||||
|
||||
|
||||
CAM_VALUE = 0.1
|
||||
CAMERA_MAP = {
|
||||
@@ -34,6 +40,256 @@ KEYBOARD_MAP_7 = { # templerun_distilled_model: still/w/s/left/right/a/d
|
||||
"d": [0, 0, 0, 0, 0, 0, 1], # d
|
||||
}
|
||||
KEYBOARD_MAP = KEYBOARD_MAP_4 # Default for backward compatibility
|
||||
SOLARIS_MOVEMENT_KEY_INDICES = {
|
||||
"W": 11,
|
||||
"S": 12,
|
||||
"A": 13,
|
||||
"D": 14,
|
||||
}
|
||||
|
||||
|
||||
def retain_kv_with_sink(
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
target_len: int,
|
||||
sink_size: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if k.shape[1] <= target_len or sink_size <= 0:
|
||||
return k[:, -target_len:], v[:, -target_len:]
|
||||
|
||||
sink_len = min(int(sink_size), target_len, k.shape[1])
|
||||
tail_len = target_len - sink_len
|
||||
sink_k = k[:, :sink_len]
|
||||
sink_v = v[:, :sink_len]
|
||||
if tail_len <= 0:
|
||||
return sink_k, sink_v
|
||||
|
||||
tail_k = k[:, sink_len:][:, -tail_len:]
|
||||
tail_v = v[:, sink_len:][:, -tail_len:]
|
||||
return torch.cat([sink_k, tail_k], dim=1), torch.cat([sink_v, tail_v], dim=1)
|
||||
|
||||
|
||||
def _require_cv2() -> bool:
|
||||
if cv2 is None:
|
||||
logger.warning(
|
||||
"OpenCV is not available; skipping MatrixGame2 validation overlay."
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _matrixgame2_overlay_keys_from_keyboard(
|
||||
keyboard_frame: np.ndarray | torch.Tensor,
|
||||
) -> dict[str, bool]:
|
||||
vec = np.asarray(keyboard_frame, dtype=np.float32).reshape(-1)
|
||||
dim = vec.shape[0]
|
||||
|
||||
if dim >= 23:
|
||||
return {
|
||||
key: bool(vec[idx] > 0.5)
|
||||
for key, idx in SOLARIS_MOVEMENT_KEY_INDICES.items()
|
||||
}
|
||||
if dim == 7:
|
||||
return {
|
||||
"W": bool(vec[1] > 0.5),
|
||||
"S": bool(vec[2] > 0.5),
|
||||
"A": bool(vec[5] > 0.5),
|
||||
"D": bool(vec[6] > 0.5),
|
||||
}
|
||||
if dim >= 6:
|
||||
return {
|
||||
"W": bool(vec[0] > 0.5),
|
||||
"S": bool(vec[1] > 0.5),
|
||||
"A": bool(vec[2] > 0.5),
|
||||
"D": bool(vec[3] > 0.5),
|
||||
}
|
||||
if dim == 4:
|
||||
return {
|
||||
"W": bool(vec[0] > 0.5),
|
||||
"S": bool(vec[1] > 0.5),
|
||||
"A": bool(vec[2] > 0.5),
|
||||
"D": bool(vec[3] > 0.5),
|
||||
}
|
||||
if dim == 2:
|
||||
return {
|
||||
"W": bool(vec[0] > 0.5),
|
||||
"S": bool(vec[1] > 0.5),
|
||||
"A": False,
|
||||
"D": False,
|
||||
}
|
||||
return {"W": False, "S": False, "A": False, "D": False}
|
||||
|
||||
|
||||
def draw_rounded_rectangle(
|
||||
image: np.ndarray,
|
||||
top_left: tuple[int, int],
|
||||
bottom_right: tuple[int, int],
|
||||
color: tuple[int, int, int],
|
||||
radius: int = 10,
|
||||
alpha: float = 0.5,
|
||||
) -> None:
|
||||
if not _require_cv2():
|
||||
return
|
||||
|
||||
overlay = image.copy()
|
||||
x1, y1 = top_left
|
||||
x2, y2 = bottom_right
|
||||
|
||||
cv2.rectangle(overlay, (x1 + radius, y1), (x2 - radius, y2), color, -1)
|
||||
cv2.rectangle(overlay, (x1, y1 + radius), (x2, y2 - radius), color, -1)
|
||||
cv2.ellipse(overlay, (x1 + radius, y1 + radius), (radius, radius), 180, 0, 90, color, -1)
|
||||
cv2.ellipse(overlay, (x2 - radius, y1 + radius), (radius, radius), 270, 0, 90, color, -1)
|
||||
cv2.ellipse(overlay, (x1 + radius, y2 - radius), (radius, radius), 90, 0, 90, color, -1)
|
||||
cv2.ellipse(overlay, (x2 - radius, y2 - radius), (radius, radius), 0, 0, 90, color, -1)
|
||||
cv2.addWeighted(overlay, alpha, image, 1 - alpha, 0, image)
|
||||
|
||||
|
||||
def draw_keys_on_frame(
|
||||
frame: np.ndarray,
|
||||
keys: dict[str, bool],
|
||||
key_size: tuple[int, int] = (30, 30),
|
||||
top_margin: int = 15,
|
||||
) -> None:
|
||||
if not _require_cv2():
|
||||
return
|
||||
|
||||
left_margin = 15
|
||||
gap = 3
|
||||
key_positions = {
|
||||
"W": (left_margin + key_size[0] + gap, top_margin),
|
||||
"A": (left_margin, top_margin + key_size[1] + gap),
|
||||
"S": (left_margin + key_size[0] + gap, top_margin + key_size[1] + gap),
|
||||
"D": (
|
||||
left_margin + (key_size[0] + gap) * 2,
|
||||
top_margin + key_size[1] + gap,
|
||||
),
|
||||
}
|
||||
|
||||
for key, (x, y) in key_positions.items():
|
||||
pressed = keys.get(key, False)
|
||||
color = (0, 255, 0) if pressed else (200, 200, 200)
|
||||
alpha = 0.8 if pressed else 0.5
|
||||
draw_rounded_rectangle(
|
||||
frame,
|
||||
(x, y),
|
||||
(x + key_size[0], y + key_size[1]),
|
||||
color,
|
||||
radius=5,
|
||||
alpha=alpha,
|
||||
)
|
||||
text_size = cv2.getTextSize(
|
||||
key, cv2.FONT_HERSHEY_SIMPLEX, 0.5, 1
|
||||
)[0]
|
||||
text_x = x + (key_size[0] - text_size[0]) // 2
|
||||
text_y = y + (key_size[1] + text_size[1]) // 2
|
||||
cv2.putText(
|
||||
frame,
|
||||
key,
|
||||
(text_x, text_y),
|
||||
cv2.FONT_HERSHEY_SIMPLEX,
|
||||
0.5,
|
||||
(0, 0, 0),
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def draw_mouse_on_frame(
|
||||
frame: np.ndarray,
|
||||
yaw: float,
|
||||
pitch: float,
|
||||
top_margin: int = 15,
|
||||
) -> None:
|
||||
if not _require_cv2():
|
||||
return
|
||||
|
||||
_, width, _ = frame.shape
|
||||
right_margin = 15
|
||||
crosshair_radius = 25
|
||||
crosshair_x = width - right_margin - crosshair_radius
|
||||
crosshair_y = top_margin + crosshair_radius
|
||||
|
||||
dx = int(yaw * crosshair_radius * 8)
|
||||
dy = int(-pitch * crosshair_radius * 8)
|
||||
max_arrow = crosshair_radius - 5
|
||||
dx = max(-max_arrow, min(max_arrow, dx))
|
||||
dy = max(-max_arrow, min(max_arrow, dy))
|
||||
|
||||
cv2.circle(frame, (crosshair_x, crosshair_y), crosshair_radius, (50, 50, 50), -1)
|
||||
cv2.circle(frame, (crosshair_x, crosshair_y), crosshair_radius, (200, 200, 200), 1)
|
||||
cv2.line(
|
||||
frame,
|
||||
(crosshair_x - crosshair_radius + 5, crosshair_y),
|
||||
(crosshair_x + crosshair_radius - 5, crosshair_y),
|
||||
(100, 100, 100),
|
||||
1,
|
||||
)
|
||||
cv2.line(
|
||||
frame,
|
||||
(crosshair_x, crosshair_y - crosshair_radius + 5),
|
||||
(crosshair_x, crosshair_y + crosshair_radius - 5),
|
||||
(100, 100, 100),
|
||||
1,
|
||||
)
|
||||
|
||||
if abs(dx) > 1 or abs(dy) > 1:
|
||||
cv2.arrowedLine(
|
||||
frame,
|
||||
(crosshair_x, crosshair_y),
|
||||
(crosshair_x + dx, crosshair_y + dy),
|
||||
(0, 255, 0),
|
||||
2,
|
||||
tipLength=0.3,
|
||||
)
|
||||
|
||||
|
||||
def overlay_validation_actions_on_frames(
|
||||
frames: list[np.ndarray],
|
||||
keyboard_cond: np.ndarray | torch.Tensor | None = None,
|
||||
mouse_cond: np.ndarray | torch.Tensor | None = None,
|
||||
) -> list[np.ndarray]:
|
||||
if (keyboard_cond is None and mouse_cond is None) or not _require_cv2():
|
||||
return frames
|
||||
|
||||
if keyboard_cond is not None and torch.is_tensor(keyboard_cond):
|
||||
keyboard_cond = keyboard_cond.detach().cpu().float().numpy()
|
||||
if mouse_cond is not None and torch.is_tensor(mouse_cond):
|
||||
mouse_cond = mouse_cond.detach().cpu().float().numpy()
|
||||
|
||||
if keyboard_cond is not None:
|
||||
keyboard_cond = np.asarray(keyboard_cond, dtype=np.float32)
|
||||
if keyboard_cond.ndim == 3 and keyboard_cond.shape[0] == 1:
|
||||
keyboard_cond = keyboard_cond[0]
|
||||
if mouse_cond is not None:
|
||||
mouse_cond = np.asarray(mouse_cond, dtype=np.float32)
|
||||
if mouse_cond.ndim == 3 and mouse_cond.shape[0] == 1:
|
||||
mouse_cond = mouse_cond[0]
|
||||
|
||||
processed_frames: list[np.ndarray] = []
|
||||
for frame_idx, frame in enumerate(frames):
|
||||
frame = np.ascontiguousarray(frame.copy())
|
||||
if keyboard_cond is not None and frame_idx < len(keyboard_cond):
|
||||
draw_keys_on_frame(
|
||||
frame,
|
||||
_matrixgame2_overlay_keys_from_keyboard(keyboard_cond[frame_idx]),
|
||||
)
|
||||
if mouse_cond is not None and frame_idx < len(mouse_cond):
|
||||
# MG action convention: mouse[0]=pitch, mouse[1]=yaw.
|
||||
pitch = float(mouse_cond[frame_idx, 0])
|
||||
yaw = float(mouse_cond[frame_idx, 1])
|
||||
draw_mouse_on_frame(frame, yaw=yaw, pitch=pitch)
|
||||
processed_frames.append(frame)
|
||||
|
||||
return processed_frames
|
||||
|
||||
|
||||
def scale_keyboard_condition(
|
||||
keyboard_condition: torch.Tensor,
|
||||
scale: float,
|
||||
) -> torch.Tensor:
|
||||
if float(scale) == 1.0:
|
||||
return keyboard_condition
|
||||
return keyboard_condition * float(scale)
|
||||
|
||||
|
||||
def expand_action_to_frames(action: dict, num_frames: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
result = {}
|
||||
|
||||
@@ -467,6 +467,18 @@ class Qwen2_5_VisionTransformerPretrainedModel(nn.Module):
|
||||
return hidden_states
|
||||
|
||||
|
||||
def _compute_default_rope_parameters(config, device=None, seq_len=None, **kwargs):
|
||||
# transformers>=5 removes the "default" entry from ROPE_INIT_FUNCTIONS and
|
||||
# moves rope_theta inside rope_parameters; replicate the 4.x default init.
|
||||
rope_params = getattr(config, "rope_parameters", None) or getattr(config, "rope_scaling", None) or {}
|
||||
base = rope_params.get("rope_theta", getattr(config, "rope_theta", 10000.0))
|
||||
head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
|
||||
partial_rotary_factor = getattr(config, "partial_rotary_factor", 1.0)
|
||||
dim = int(head_dim * partial_rotary_factor)
|
||||
inv_freq = 1.0 / (base**(torch.arange(0, dim, 2, dtype=torch.int64).float().to(device) / dim))
|
||||
return inv_freq, 1.0
|
||||
|
||||
|
||||
class Qwen2_5_VLRotaryEmbedding(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, device=None):
|
||||
super().__init__()
|
||||
@@ -479,7 +491,14 @@ class Qwen2_5_VLRotaryEmbedding(nn.Module):
|
||||
self.original_max_seq_len = config.max_position_embeddings
|
||||
|
||||
self.config = config
|
||||
self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
|
||||
if self.rope_type in ROPE_INIT_FUNCTIONS:
|
||||
self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
|
||||
elif self.rope_type == "default":
|
||||
# transformers>=5 drops the "default" entry from ROPE_INIT_FUNCTIONS.
|
||||
self.rope_init_fn = _compute_default_rope_parameters
|
||||
else:
|
||||
raise KeyError(f"Unsupported rope_type '{self.rope_type}'; available: "
|
||||
f"{['default', *ROPE_INIT_FUNCTIONS]}")
|
||||
|
||||
inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
@@ -948,6 +967,16 @@ QWEN2_5_VL_ATTENTION_CLASSES = {
|
||||
# If FlashAttention2 is not available, transparently fall back to SDPA.
|
||||
if not is_flash_attn_2_available():
|
||||
QWEN2_5_VL_ATTENTION_CLASSES["flash_attention_2"] = Qwen2_5_VLSdpaAttention
|
||||
else:
|
||||
# transformers>=5 only resolves the flash-attn functions when the model
|
||||
# preloads them via its attention interface; this module bypasses
|
||||
# PreTrainedModel, so _flash_attention_forward(implementation=None) raises
|
||||
# unless we preload here.
|
||||
try:
|
||||
from transformers.modeling_flash_attention_utils import lazy_import_flash_attention
|
||||
lazy_import_flash_attention("flash_attention_2")
|
||||
except (ImportError, ValueError):
|
||||
QWEN2_5_VL_ATTENTION_CLASSES["flash_attention_2"] = Qwen2_5_VLSdpaAttention
|
||||
|
||||
class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int):
|
||||
@@ -1035,7 +1064,7 @@ class Qwen2_5_VLModel(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.padding_idx = getattr(config, "pad_token_id", None)
|
||||
self.vocab_size = config.vocab_size
|
||||
|
||||
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
|
||||
@@ -1408,6 +1437,18 @@ class Qwen2_5_VLCausalLMOutputWithPast(ModelOutput):
|
||||
rope_deltas: Optional[torch.LongTensor] = None
|
||||
|
||||
|
||||
def _flatten_text_config(config):
|
||||
# transformers>=5 stops forwarding text-model attributes (hidden_size,
|
||||
# vocab_size, rope_scaling, ...) from the composite Qwen2_5_VLConfig to
|
||||
# config.text_config; this module reads them from the top level.
|
||||
text_config = getattr(config, "text_config", None)
|
||||
if text_config is not None:
|
||||
for key, value in text_config.to_dict().items():
|
||||
if not hasattr(config, key):
|
||||
setattr(config, key, value)
|
||||
return config
|
||||
|
||||
|
||||
class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
|
||||
_tied_weights_keys = ["lm_head.weight"]
|
||||
config_class = Qwen2_5_VLConfig
|
||||
@@ -1415,6 +1456,7 @@ class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
config = _flatten_text_config(config)
|
||||
self.config = config
|
||||
self.visual = Qwen2_5_VisionTransformerPretrainedModel(config.vision_config)
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user