Compare commits
82
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4e1603634d | ||
|
|
f1eeb6303f | ||
|
|
990d2c2410 | ||
|
|
1772226bf1 | ||
|
|
1fe8e64e23 | ||
|
|
46f18d8cd2 | ||
|
|
2e8db18d94 | ||
|
|
4190c7203f | ||
|
|
8f1443f47b | ||
|
|
42ed546a66 | ||
|
|
d6e020402e | ||
|
|
620e100af4 | ||
|
|
0fde316a19 | ||
|
|
0d99e47e16 | ||
|
|
27f6f0aacd | ||
|
|
c97fb6b3b3 | ||
|
|
3464cb8b03 | ||
|
|
e114fba53f | ||
|
|
093f5e699c | ||
|
|
e24bc12c59 | ||
|
|
4d04c1b01c | ||
|
|
3fb150bbe0 | ||
|
|
320f8a1f8d | ||
|
|
ccfcc3042b | ||
|
|
8f0493637e | ||
|
|
f5ce12f17a | ||
|
|
17cb6737c1 | ||
|
|
f39dbe482c | ||
|
|
5b5608cb37 | ||
|
|
1d4a6037eb | ||
|
|
8ac6526cdc | ||
|
|
08364b2080 | ||
|
|
f81de3926f | ||
|
|
04633096e2 | ||
|
|
b6be3d0c8a | ||
|
|
d5acc7bbae | ||
|
|
9035927da1 | ||
|
|
42d1c79694 | ||
|
|
59000cb933 | ||
|
|
e6b15fc4de | ||
|
|
818daea816 | ||
|
|
d1bd1d8da4 | ||
|
|
d210a076e6 | ||
|
|
7e96218d04 | ||
|
|
6d0ddb44bc | ||
|
|
14c142ba9c | ||
|
|
3fa3f4e333 | ||
|
|
5c7cd391ac | ||
|
|
e7d6c40860 | ||
|
|
12a7cb53ff | ||
|
|
c17d33bf33 | ||
|
|
2aaeee2ab8 | ||
|
|
eb3a394224 | ||
|
|
f673423b51 | ||
|
|
eb0a41528a | ||
|
|
140bd1a6cf | ||
|
|
11f5a8e582 | ||
|
|
71b3cb8c34 | ||
|
|
c85f6a477f | ||
|
|
40d4930d73 | ||
|
|
f9be085243 | ||
|
|
36b53ff350 | ||
|
|
9801037c3d | ||
|
|
74d09b0efd | ||
|
|
38dc8820ac | ||
|
|
c77a76c6af | ||
|
|
d14d5aadea | ||
|
|
4c915b7742 | ||
|
|
9a8bbe18fa | ||
|
|
ea25441ef0 | ||
|
|
48957fcde1 | ||
|
|
7b872cc41e | ||
|
|
37418946c8 | ||
|
|
95fd29e0cb | ||
|
|
e17cd2633c | ||
|
|
e0dc5f2b0c | ||
|
|
70ee5d230c | ||
|
|
24ced500f5 | ||
|
|
4ddcdf541f | ||
|
|
0e3529869c | ||
|
|
e1e0d91c00 | ||
|
|
145a3f166b |
@@ -119,7 +119,7 @@ FastVideo-WorldModel/
|
||||
## Build & Test Commands
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev] # Editable install
|
||||
uv pip install -e ".[dev]" # Editable install
|
||||
pre-commit run --all-files # Lint/format/spell
|
||||
pytest tests/ # Top-level tests
|
||||
pytest fastvideo/tests/ -v # Package tests
|
||||
|
||||
Executable
+96
@@ -0,0 +1,96 @@
|
||||
#!/usr/bin/env bash
|
||||
# Sync .agents/skills/ into .claude/skills/ via per-skill symlinks.
|
||||
#
|
||||
# Why: Claude Code only scans .claude/skills/ and ~/.claude/skills/ for
|
||||
# user-invocable skills (no skillsPath config exists — see
|
||||
# https://code.claude.com/docs/en/skills.md). This repo's skills live
|
||||
# in .agents/skills/ so they travel with the repo and stay under git.
|
||||
# Run this once after cloning (or after adding/removing a skill) to
|
||||
# expose them to Claude Code without maintaining a parallel tree.
|
||||
#
|
||||
# Usage:
|
||||
# .agents/scripts/sync-skills.sh
|
||||
#
|
||||
# Idempotent and safe to re-run. Prunes stale symlinks whose source
|
||||
# has been removed from .agents/skills/. Leaves hand-written
|
||||
# .claude/skills/<name>/ directories untouched (only symlinks are
|
||||
# managed).
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(git -C "$(dirname "$0")" rev-parse --show-toplevel)"
|
||||
SRC_DIR="$REPO_ROOT/.agents/skills"
|
||||
DST_DIR="$REPO_ROOT/.claude/skills"
|
||||
|
||||
if [[ ! -d "$SRC_DIR" ]]; then
|
||||
echo "Error: $SRC_DIR does not exist." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir -p "$DST_DIR"
|
||||
|
||||
linked=0
|
||||
unchanged=0
|
||||
skipped=0
|
||||
pruned=0
|
||||
|
||||
link_skill() {
|
||||
local name="$1"
|
||||
local src="$SRC_DIR/$name"
|
||||
local dst="$DST_DIR/$name"
|
||||
# Relative target keeps symlinks portable across clones.
|
||||
local rel="../../.agents/skills/$name"
|
||||
|
||||
if [[ -L "$dst" ]]; then
|
||||
if [[ "$(readlink "$dst")" == "$rel" ]]; then
|
||||
unchanged=$((unchanged + 1))
|
||||
return
|
||||
fi
|
||||
rm "$dst"
|
||||
elif [[ -e "$dst" ]]; then
|
||||
echo "Skipped (not a symlink): .claude/skills/$name" >&2
|
||||
skipped=$((skipped + 1))
|
||||
return
|
||||
fi
|
||||
|
||||
ln -s "$rel" "$dst"
|
||||
echo "Linked: .claude/skills/$name -> $rel"
|
||||
linked=$((linked + 1))
|
||||
}
|
||||
|
||||
prune_stale() {
|
||||
local link="$1"
|
||||
local target
|
||||
target="$(readlink "$link")"
|
||||
case "$target" in
|
||||
../../.agents/skills/*) ;;
|
||||
*) return ;;
|
||||
esac
|
||||
local name="${target##*/}"
|
||||
if [[ ! -d "$SRC_DIR/$name" ]]; then
|
||||
rm "$link"
|
||||
echo "Pruned stale: .claude/skills/$(basename "$link")"
|
||||
pruned=$((pruned + 1))
|
||||
fi
|
||||
}
|
||||
|
||||
for src in "$SRC_DIR"/*/; do
|
||||
[[ -d "$src" ]] || continue
|
||||
name="$(basename "$src")"
|
||||
# Only treat directories that actually contain a SKILL.md as skills.
|
||||
[[ -f "$src/SKILL.md" ]] || continue
|
||||
link_skill "$name"
|
||||
done
|
||||
|
||||
shopt -s nullglob
|
||||
for link in "$DST_DIR"/*; do
|
||||
[[ -L "$link" ]] || continue
|
||||
prune_stale "$link"
|
||||
done
|
||||
shopt -u nullglob
|
||||
|
||||
printf "\nSummary: %d linked, %d unchanged, %d pruned" "$linked" "$unchanged" "$pruned"
|
||||
if [[ "$skipped" -gt 0 ]]; then
|
||||
printf ", %d skipped (non-symlink collision)" "$skipped"
|
||||
fi
|
||||
printf "\n"
|
||||
@@ -5,3 +5,6 @@
|
||||
{"name": "evaluate-video-quality", "description": "Evaluate generated video quality using available metrics (SSIM, loss trajectory, caption consistency)", "path": "evaluate-video-quality/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "reseed-ssim-references", "description": "Re-seed (overwrite) HF reference videos for an existing fastvideo/tests/ssim/ test and a single model id on Modal L40S. Always backs up current refs first, regenerates on Modal, pauses for the user to eyeball before-vs-after, then uploads with --force scoped to --model-id. Sister skill to seed-ssim-references; use when intentional code change has invalidated existing refs", "path": "reseed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "add-model", "description": "Add a new model (or variant) to FastVideo: DiT + configs + pipeline + presets + registry + tests. Walks through FastVideo's single stage-based pipeline architecture with exact file paths and registration hooks.", "path": "add-model/SKILL.md", "status": "draft", "trust": "low"}
|
||||
|
||||
@@ -12,7 +12,7 @@ automates the boilerplate of setting environment variables, picking the right
|
||||
entrypoint, and applying defaults from the closest example script.
|
||||
|
||||
## Prerequisites
|
||||
- The repo is cloned and `fastvideo` is installed (`uv pip install -e .[dev]`).
|
||||
- The repo is cloned and `fastvideo` is installed (`uv pip install -e ".[dev]"`).
|
||||
- Dataset is preprocessed (see `docs/training/data_preprocess.md`).
|
||||
- `WANDB_API_KEY` is set in the environment (or `WANDB_MODE=offline` for local).
|
||||
- GPU resources are available (multi-GPU requires NCCL).
|
||||
|
||||
@@ -0,0 +1,343 @@
|
||||
---
|
||||
name: reseed-ssim-references
|
||||
description: Re-seed HF reference videos for a single existing SSIM test on Modal L40S. Always backs up current refs locally first, regenerates on Modal, pauses for the user to eyeball before-vs-after quality, then overwrites the targeted `<model_id>` subtree on `FastVideo/ssim-reference-videos` with `--force`. Use when an intentional code change (model port fix, attention backend swap, kernel upgrade, hyperparameter change) has invalidated existing refs and they need to be regenerated. Pairs with `seed-ssim-references`, which is for first-time seeding only.
|
||||
---
|
||||
|
||||
# Re-seed SSIM Reference Videos
|
||||
|
||||
## Purpose
|
||||
|
||||
Replace the existing SSIM reference videos for a single `(test_file, model_id)`
|
||||
pair on the HF dataset (`FastVideo/ssim-reference-videos`). This is **destructive**
|
||||
on HF — the old refs are overwritten — so the skill always:
|
||||
|
||||
1. Confirms intent with a one-liner the user has to type.
|
||||
2. Downloads the existing refs as a local, timestamped backup.
|
||||
3. Regenerates on Modal L40S (same code path that CI uses).
|
||||
4. Pauses for a side-by-side eyeball of backup vs new mp4s.
|
||||
5. Uploads with `--force`, scoped to the single `--model-id`.
|
||||
6. Reminds the user to keep the backup until the PR lands.
|
||||
|
||||
Pairs with `seed-ssim-references`, which is the inverse (first-time seeding
|
||||
only, refuses to overwrite). Re-seeding is intentionally a separate, more
|
||||
ceremonial operation because mistakenly clobbering production refs is much
|
||||
harder to recover from than failing closed.
|
||||
|
||||
## When to use
|
||||
|
||||
- An intentional code change (model port fix, kernel upgrade, attention
|
||||
backend swap, hyperparameter change in the test itself) has shifted the
|
||||
expected SSIM output and the existing refs no longer represent the new
|
||||
ground truth.
|
||||
- A test is failing in CI **for the right reason** (the new code is correct,
|
||||
the old refs are stale).
|
||||
|
||||
## When not to use
|
||||
|
||||
- A test is failing for the **wrong** reason (the port is buggy, not the
|
||||
refs). Fix the port; re-seeding hides the bug.
|
||||
- A brand-new test that has no refs on HF yet. Use `seed-ssim-references`.
|
||||
- "Just to clean up drift" without a concrete code change to point at. The
|
||||
PR description has to justify *why* refs changed; without a concrete
|
||||
change, there's nothing to write.
|
||||
|
||||
## Inputs
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `test_file` | Yes | Path to the SSIM test, e.g. `fastvideo/tests/ssim/test_matrixgame_similarity.py`. Validated against `fastvideo/tests/ssim/test_*_similarity.py`. |
|
||||
| `model_id` | Yes | Single model id from the test's `*_MODEL_TO_PARAMS`, e.g. `Matrix-Game-2.0-Diffusers-Base`. Re-seed runs are **per model**. For multi-model tests, invoke the skill once per model. |
|
||||
| `intent_rationale` | Yes | One-line explanation of *why* refs are being regenerated (e.g. "Relax FA-2 head_size whitelist to include 80 — matrix_game now uses FLASH_ATTN instead of TORCH_SDPA"). Recorded in the backup directory and reused in the PR description. |
|
||||
|
||||
Hardcoded:
|
||||
|
||||
- Modal GPU: **L40S** (matches CI; re-seeding from another SKU produces refs
|
||||
that L40S CI cannot match).
|
||||
- Quality tier: **`default`**. `full_quality` is a separate, deliberate
|
||||
operation.
|
||||
- HF repo: `FastVideo/ssim-reference-videos` (override via
|
||||
`FASTVIDEO_SSIM_REFERENCE_HF_REPO`).
|
||||
- Device folder: `L40S_reference_videos`.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
The user has confirmed:
|
||||
|
||||
- `modal` CLI authenticated.
|
||||
- `hf` CLI authenticated, **and** `HF_API_KEY` (or `HUGGINGFACE_HUB_TOKEN` /
|
||||
`HF_TOKEN`) exported with **write** access to
|
||||
`FastVideo/ssim-reference-videos`.
|
||||
- The current branch's code is the change that motivated the re-seed (i.e.
|
||||
`git rev-parse HEAD` is the commit that intentionally invalidated refs).
|
||||
|
||||
Fail fast if any of these are missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Validate inputs and confirm intent
|
||||
|
||||
- Verify `test_file` exists and matches `fastvideo/tests/ssim/test_*_similarity.py`.
|
||||
- Grep the file for `*_MODEL_TO_PARAMS` and assert `model_id` is one of its
|
||||
keys. If the file has only a single hardcoded model, accept that model id
|
||||
as the only valid value.
|
||||
- Print the rationale and ask the user to type **`confirm reseed`** (not just
|
||||
`y` — make it deliberate):
|
||||
|
||||
> About to RE-SEED references for model `<model_id>` from test `<test_file>`.
|
||||
> This will OVERWRITE existing refs on
|
||||
> `FastVideo/ssim-reference-videos/reference_videos/default/L40S_reference_videos/<model_id>/`
|
||||
> after backup + Modal regen + eyeball.
|
||||
>
|
||||
> Reason: `<intent_rationale>`
|
||||
> HEAD: `<git rev-parse --short=12 HEAD>`
|
||||
>
|
||||
> Reply `confirm reseed` to proceed, anything else to abort.
|
||||
|
||||
Stop until the user types exactly `confirm reseed`. Anything else aborts
|
||||
with no side effects.
|
||||
|
||||
### 2. Back up existing refs
|
||||
|
||||
Always required. The backup is the only graceful path back if anything goes
|
||||
wrong later.
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(echo "<model_id>" | tr '/' '_')
|
||||
BACKUP_DIR="ssim_reseed_backup/${TIMESTAMP}_${SHORT_COMMIT}_${MODEL_SAFE}"
|
||||
mkdir -p "$BACKUP_DIR"
|
||||
|
||||
hf download \
|
||||
--repo-type dataset FastVideo/ssim-reference-videos \
|
||||
--include "reference_videos/default/L40S_reference_videos/<model_id>/**" \
|
||||
--local-dir "$BACKUP_DIR"
|
||||
|
||||
mp4_count=$(find "$BACKUP_DIR" -name "*.mp4" | wc -l)
|
||||
echo "Backup mp4 count: $mp4_count"
|
||||
[ "$mp4_count" -gt 0 ] || {
|
||||
echo "ERROR: backup is empty for <model_id>. Either the model id is wrong"
|
||||
echo "or there are no existing refs (use seed-ssim-references instead)."
|
||||
exit 1
|
||||
}
|
||||
|
||||
# Provenance — used in the PR description
|
||||
cat > "$BACKUP_DIR/PROVENANCE.txt" <<EOF
|
||||
test_file: <test_file>
|
||||
model_id: <model_id>
|
||||
head_commit: $(git rev-parse HEAD)
|
||||
timestamp_utc: $(date -u +%FT%TZ)
|
||||
reason: <intent_rationale>
|
||||
EOF
|
||||
```
|
||||
|
||||
If the `hf download` produces zero mp4s, abort — the user has either picked a
|
||||
non-existent `model_id` or there are no refs yet (in which case
|
||||
`seed-ssim-references` is the right tool).
|
||||
|
||||
### 3. Regenerate on Modal L40S
|
||||
|
||||
Mirror CI's exact env recipe so the regenerated refs are byte-comparable to
|
||||
what CI will produce on the same commit. Two differences from CI:
|
||||
|
||||
1. **Pass the same env prefix CI uses** (`IMAGE_VERSION`, `BUILDKITE_*`) — see
|
||||
`.buildkite/pipeline.yml:1-3` and `.buildkite/scripts/pr_test.sh:62-83`.
|
||||
Without this, `ssim_test.py:17-18` resolves a different GHCR image tag
|
||||
(default is `latest`, CI is `py3.12-latest`), and `ssim_test.py:38-46`
|
||||
bakes different values into the image's frozen env block. **Mismatched
|
||||
image or env is the most common source of SSIM drift between reseed and
|
||||
CI runs.**
|
||||
2. **Do not pass `--skip-reference-download`**. Letting the test fetch the
|
||||
existing refs and run the full SSIM compare gives "before" SSIM numbers
|
||||
for the PR description, and the test still produces the new mp4s
|
||||
regardless of whether the comparison passes or fails.
|
||||
|
||||
```bash
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
|
||||
IMAGE_VERSION="py3.12-latest" \
|
||||
BUILDKITE_REPO="$(git config --get remote.origin.url)" \
|
||||
BUILDKITE_COMMIT="$(git rev-parse HEAD)" \
|
||||
BUILDKITE_PULL_REQUEST="${BUILDKITE_PULL_REQUEST:-false}" \
|
||||
modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--git-repo="$(git config --get remote.origin.url)" \
|
||||
--git-commit="$(git rev-parse HEAD)" \
|
||||
--hf-api-key="$HF_API_KEY" \
|
||||
--test-files="<test_file>" \
|
||||
--sync-generated-to-volume \
|
||||
--generated-volume-subdir="$SUBDIR" \
|
||||
--no-fail-fast
|
||||
```
|
||||
|
||||
Capture the printed `modal volume get ...` hint — its `<SUBDIR>` matches
|
||||
`$SUBDIR` and is needed for step 4. Capture the SSIM numbers from the test
|
||||
output (or from the JSON next to the generated mp4) for the PR description.
|
||||
|
||||
### 4. Download generated videos
|
||||
|
||||
```bash
|
||||
modal volume get --force hf-model-weights \
|
||||
ssim_generated_videos/default/"$SUBDIR"/generated_videos \
|
||||
./generated_videos_modal/default
|
||||
```
|
||||
|
||||
After this, the new mp4s live at:
|
||||
|
||||
```
|
||||
./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4
|
||||
```
|
||||
|
||||
`--force` is required when `./generated_videos_modal/default` already exists
|
||||
from a prior run; safe on the first run too.
|
||||
|
||||
### 5. PAUSE — user reviews quality side-by-side
|
||||
|
||||
Print the diff and the comparison:
|
||||
|
||||
```bash
|
||||
echo "=== File list diff (backup vs new) ==="
|
||||
diff -u \
|
||||
<(find "$BACKUP_DIR/reference_videos/default/L40S_reference_videos/<model_id>" -name "*.mp4" \
|
||||
| sed "s|$BACKUP_DIR/reference_videos/default/L40S_reference_videos/||" | sort) \
|
||||
<(find ./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id> -name "*.mp4" \
|
||||
| sed "s|./generated_videos_modal/default/generated_videos/L40S_reference_videos/||" | sort) \
|
||||
|| true
|
||||
|
||||
echo
|
||||
echo "=== SSIM numbers from this run (paste into PR) ==="
|
||||
find ./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id> -name "*_ssim.json" -exec cat {} \;
|
||||
```
|
||||
|
||||
Then stop and tell the user:
|
||||
|
||||
> Old refs backed up to `$BACKUP_DIR`.
|
||||
> New videos in `./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/`.
|
||||
>
|
||||
> Open both in a video player. Confirm the new videos:
|
||||
> 1. Look correct (no obvious artifacts, no black/static frames).
|
||||
> 2. Are *intentionally* different from the backup in the way described
|
||||
> in `<intent_rationale>` (e.g. slight numerical drift only, not a
|
||||
> different scene / different motion / corrupted output).
|
||||
>
|
||||
> Reply **`upload`** to overwrite HF, anything else to abort.
|
||||
> Aborting leaves the backup and new videos on disk for inspection — nothing
|
||||
> on HF changes.
|
||||
|
||||
Do not proceed until the user types exactly `upload`. If they abort, leave
|
||||
everything on disk and stop here.
|
||||
|
||||
### 6. Copy into the local reference layout
|
||||
|
||||
Same as `seed-ssim-references` step 5:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--generated-dir ./generated_videos_modal/default/generated_videos/L40S_reference_videos
|
||||
```
|
||||
|
||||
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
|
||||
### 7. Upload with `--force`, scoped to `--model-id`
|
||||
|
||||
The `--force` flag is what makes this skill different from `seed-ssim-references`.
|
||||
Always pair it with `--model-id` so a typo cannot accidentally overwrite a
|
||||
neighboring model's refs.
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>" \
|
||||
--force
|
||||
```
|
||||
|
||||
The CLI's overwrite guard refuses without `--force`; with `--force` it
|
||||
overwrites only files under
|
||||
`reference_videos/default/L40S_reference_videos/<model_id>/`.
|
||||
|
||||
### 8. Report success and retention guidance
|
||||
|
||||
Print:
|
||||
|
||||
- The HF path that was overwritten (`<repo>/reference_videos/default/L40S_reference_videos/<model_id>/`).
|
||||
- The local backup directory path.
|
||||
- The new SSIM numbers from step 5.
|
||||
- This restore command, in case the PR review surfaces a problem after
|
||||
upload:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>" \
|
||||
--reference-dir "$BACKUP_DIR/reference_videos/default/L40S_reference_videos" \
|
||||
--force
|
||||
```
|
||||
|
||||
- This PR-description checklist (see `fastvideo/tests/ssim/AGENTS.md` →
|
||||
*Updating Reference Videos*):
|
||||
1. Source commit that produced the new refs (HEAD at re-seed time).
|
||||
2. Test command and GPU SKU (`L40S`).
|
||||
3. Before/after SSIM numbers.
|
||||
4. The `<intent_rationale>` from step 1.
|
||||
5. A note that the backup lives at `$BACKUP_DIR` and should be retained
|
||||
until CI on the PR is green.
|
||||
|
||||
Do **not** auto-rerun the SSIM test — the user does that as part of the PR.
|
||||
|
||||
## Failure modes and how to handle them
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before step 2.
|
||||
- **Backup is empty (zero mp4s).** Stop before step 3 — the model id is
|
||||
wrong or the refs don't exist yet (use `seed-ssim-references`).
|
||||
- **Modal run fails before generation.** No mp4s on the volume. Don't
|
||||
upload. Investigate the failure (test crash, OOM, partition exhaustion),
|
||||
fix, then retry from step 3. Backup is still intact.
|
||||
- **Quality regressed (visual or metric).** User aborts at step 5. Backup
|
||||
retained. New videos retained on disk for inspection. Nothing on HF
|
||||
changed. Either fix the underlying code change or abandon the re-seed.
|
||||
- **User confirmed `upload` but later realized the new refs are wrong.**
|
||||
Run the restore command from step 8 with the backup `--reference-dir`.
|
||||
This is exactly why the backup exists.
|
||||
- **Multi-model test, only one model is being re-seeded.** Run the skill
|
||||
once per model id. The `--model-id` scope on upload guarantees the others
|
||||
are untouched.
|
||||
|
||||
## Design notes (for future skill maintainers)
|
||||
|
||||
- Per-`model_id` scope is mandatory. The dataset houses many model subtrees;
|
||||
re-seeding the wrong one is hard to undo without backup.
|
||||
- `default` tier only; `full_quality` is a separate, deliberate operation
|
||||
with different params and ~doubled runtime, and isn't what CI gates on.
|
||||
- The skill deliberately does **not** pass `--skip-reference-download` to
|
||||
Modal so we get pre-reseed SSIM numbers for the PR. The `seed`-skill
|
||||
passes it because no refs exist yet; for re-seed, refs do exist and
|
||||
exposing the comparison is informative.
|
||||
- The two-token confirm (`confirm reseed`, then `upload`) is intentional.
|
||||
Re-seeding is high-blast-radius and should not be one-keystroke.
|
||||
- The backup directory is plain mp4s + `PROVENANCE.txt`. No HF metadata is
|
||||
preserved; the restore path uses `reference_videos_cli.py upload
|
||||
--reference-dir` which doesn't need it.
|
||||
|
||||
## References
|
||||
|
||||
- `.agents/skills/seed-ssim-references/SKILL.md` — the first-time seed
|
||||
skill this one parallels. Read it for the Modal flag rationale shared
|
||||
between the two flows.
|
||||
- `fastvideo/tests/ssim/AGENTS.md` — directory rules, including the PR
|
||||
expectations for any reference-video change (rationale, before/after
|
||||
SSIM, source commit/model/backend).
|
||||
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
|
||||
(with `--model-id`, `--force`), `download`. The overwrite guard at
|
||||
`upload_reference_videos` is the safety net this skill leans on.
|
||||
- `fastvideo/tests/modal/ssim_test.py` — Modal orchestrator;
|
||||
`--sync-generated-to-volume`, `--generated-volume-subdir`,
|
||||
`--skip-reference-download`, `--no-fail-fast`.
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|------|--------|
|
||||
| 2026-05-02 | Initial version. Sister skill to `seed-ssim-references`, scoped to single `(test_file, model_id)` re-seeds, with mandatory backup and two-token confirm. |
|
||||
@@ -0,0 +1,376 @@
|
||||
---
|
||||
name: seed-ssim-references
|
||||
description: Seed HF reference artefacts for a single newly-added SSIM test (pixel `.mp4` for `run_text_to_video_similarity_test`-style tests, or latent `.pt` for `run_text_to_latent_similarity_test`-style tests). Runs the test on Modal L40S, downloads the generated artefacts via `modal volume get`, pauses for the user to verify (visual eyeball for mp4, numerics dump for pt), then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
|
||||
---
|
||||
|
||||
# Seed SSIM Reference Artefacts (mp4 or pt)
|
||||
|
||||
## Purpose
|
||||
|
||||
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
|
||||
reference artefacts exist on the HF dataset
|
||||
(`FastVideo/ssim-reference-videos`). The dataset hosts two kinds of artefacts
|
||||
side-by-side per `(model_id, backend, prompt)`:
|
||||
|
||||
- **`.mp4`** — pixel ground-truth for tests that call
|
||||
`run_text_to_video_similarity_test` / `run_image_to_video_similarity_test`
|
||||
in `inference_similarity_utils.py`. Compared via SSIM.
|
||||
- **`.pt`** — pre-VAE latent bundle (fp16 full latent + fp32 slice +
|
||||
metadata + `slice_spec` + `format_version`) for tests that call
|
||||
`run_text_to_latent_similarity_test` in `latent_similarity_utils.py`.
|
||||
Compared via cosine distance on the slice and the full tensor.
|
||||
|
||||
This skill:
|
||||
|
||||
1. Detects which artefact type the test produces (pixel vs latent).
|
||||
2. Runs the test on Modal's L40S pool to generate the artefacts.
|
||||
3. Downloads them to the local repo via `modal volume get`.
|
||||
4. Pauses so the user can verify quality:
|
||||
- **mp4**: visual eyeball in a video player.
|
||||
- **pt**: numerics dump (shape, slice stats, NaN/Inf check, metadata).
|
||||
5. Uploads only the new test's files to HF, with a guard that refuses to
|
||||
overwrite anything already present.
|
||||
|
||||
The skill is run **manually**, once per new test. Before invoking it, the user
|
||||
has already sanity-tested the new test locally — it launches `VideoGenerator`
|
||||
and writes an artefact without crashing (the missing-reference assertion at
|
||||
the end is expected). The skill does not re-test locally; it goes straight
|
||||
to Modal L40S (which is what CI uses).
|
||||
|
||||
## When to use
|
||||
|
||||
- A new `test_*_similarity.py` file has been added in `fastvideo/tests/ssim/`
|
||||
and the HF dataset has no `reference_videos/default/L40S_reference_videos/<model_id>/`
|
||||
subtree for it yet.
|
||||
|
||||
## When not to use
|
||||
|
||||
- Regular CI runs — once refs exist, `pytest fastvideo/tests/ssim/` downloads
|
||||
them automatically.
|
||||
- Re-seeding an existing test. That requires `--force` on the upload step, and
|
||||
is out of scope here; treat as a separate, deliberate operation.
|
||||
|
||||
## Inputs
|
||||
|
||||
The skill has **one required input**: the path to the new SSIM test file.
|
||||
Prompt the user for it if they didn't supply it.
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `test_file` | Yes | e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`. The skill's first action is to ask for this if missing. |
|
||||
|
||||
Everything else is fixed:
|
||||
|
||||
- Modal runner GPU: **L40S** (hardcoded in `fastvideo/tests/modal/ssim_test.py`).
|
||||
- Device folder: `L40S_reference_videos`.
|
||||
- Quality tier: `default` (the tier CI runs). The `full_quality` tier is not
|
||||
seeded by this skill.
|
||||
- HF repo: `FastVideo/ssim-reference-videos` (dataset).
|
||||
- Multi-model test files: all model ids in `*_MODEL_TO_PARAMS` are seeded
|
||||
together; the Modal run produces one mp4 per (model, prompt, backend) and
|
||||
the upload scopes by `--model-id`, looping if there is more than one.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
The user has confirmed:
|
||||
|
||||
- `modal` CLI authenticated.
|
||||
- `HF_API_KEY` (or `HUGGINGFACE_HUB_TOKEN` / `HF_TOKEN`) exported with write
|
||||
access to `FastVideo/ssim-reference-videos`.
|
||||
- The test file runs locally end-to-end (generates an mp4; SSIM assertion
|
||||
failure due to missing reference is expected and fine).
|
||||
|
||||
Fail fast if the token env var is missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Ask for the test file, then detect artefact type
|
||||
|
||||
If the user didn't name one, ask: *"Which SSIM test file do you want to seed
|
||||
references for? (e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`)"*.
|
||||
|
||||
Validate:
|
||||
|
||||
- Path exists and matches `fastvideo/tests/ssim/test_*_similarity.py`.
|
||||
- File defines a `*_MODEL_TO_PARAMS` dict — grep it to extract the set of
|
||||
model ids. Those ids drive step 5.
|
||||
|
||||
Detect artefact type by inspecting the file's imports / helper call:
|
||||
|
||||
- **latent** (`.pt`) — file imports `run_text_to_latent_similarity_test`
|
||||
from `fastvideo.tests.ssim.latent_similarity_utils` (or any other helper
|
||||
that ends with `_latent_similarity_test`).
|
||||
- **pixel** (`.mp4`) — file imports
|
||||
`run_text_to_video_similarity_test` / `run_image_to_video_similarity_test`
|
||||
from `fastvideo.tests.ssim.inference_similarity_utils`, OR uses the
|
||||
legacy custom-inline helper pattern (see `test_gamecraft`,
|
||||
`test_longcat`, etc.). Default to pixel when both heuristics fail.
|
||||
|
||||
Record `ARTEFACT_TYPE ∈ {pixel, latent}` for use in step 4. Steps 2, 3, 5,
|
||||
and 6 are artefact-type-agnostic — `_iter_reference_files`,
|
||||
`copy_generated_to_reference`, and `upload_reference_videos` already walk
|
||||
both `.mp4` and `.pt` (see `reference_videos_cli.py`).
|
||||
|
||||
If either check fails, stop and tell the user what's wrong.
|
||||
|
||||
### 2. Run the test on Modal L40S
|
||||
|
||||
Pick a subdir name so repeated runs don't collide:
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
```
|
||||
|
||||
Then launch the Modal run. The `IMAGE_VERSION` and `BUILDKITE_*` env-prefix
|
||||
**must** match what CI exports in `.buildkite/scripts/pr_test.sh`, otherwise
|
||||
`fastvideo/tests/modal/ssim_test.py` resolves a different GHCR image tag
|
||||
(default is `latest`, CI is `py3.12-latest`) and bakes different values into
|
||||
the image's frozen env block (`ssim_test.py:17-18, 38-46`). Mismatched image
|
||||
or env produces SSIM drift that doesn't show up until the same commit runs
|
||||
in CI.
|
||||
|
||||
```bash
|
||||
IMAGE_VERSION="py3.12-latest" \
|
||||
BUILDKITE_REPO="$(git config --get remote.origin.url)" \
|
||||
BUILDKITE_COMMIT="$(git rev-parse HEAD)" \
|
||||
BUILDKITE_PULL_REQUEST="${BUILDKITE_PULL_REQUEST:-false}" \
|
||||
modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--git-repo="$(git config --get remote.origin.url)" \
|
||||
--git-commit="$(git rev-parse HEAD)" \
|
||||
--hf-api-key="$HF_API_KEY" \
|
||||
--test-files="<test_file>" \
|
||||
--sync-generated-to-volume \
|
||||
--generated-volume-subdir="$SUBDIR" \
|
||||
--skip-reference-download \
|
||||
--no-fail-fast
|
||||
```
|
||||
|
||||
Env prefix rationale (parity with CI; see `.buildkite/pipeline.yml:1-3` and
|
||||
`.buildkite/scripts/pr_test.sh:62-83`):
|
||||
- `IMAGE_VERSION=py3.12-latest`: pins the Modal image tag to the same one CI
|
||||
uses. Without this, `ssim_test.py:17` falls back to `latest`, which on
|
||||
GHCR is built from `Dockerfile.python3.10` — different Python, torch, and
|
||||
flash-attn wheel than CI's `py3.12-latest` (`infra-build-image.yml:51-67`,
|
||||
`_template-build-image.yml:65-101`).
|
||||
- `BUILDKITE_REPO`/`BUILDKITE_COMMIT`/`BUILDKITE_PULL_REQUEST`: mirror what
|
||||
Buildkite exports. `ssim_test.py:38-46` bakes these into the image's
|
||||
`.env(...)` block; mismatched values can perturb in-container code paths
|
||||
that branch on PR-vs-non-PR. `false` for `BUILDKITE_PULL_REQUEST` matches
|
||||
Buildkite's "non-PR build" sentinel.
|
||||
|
||||
Flag rationale:
|
||||
- `--skip-reference-download`: no refs exist yet, so conftest must not try to
|
||||
pull them.
|
||||
- `--no-fail-fast`: lets the test finish generation before `_assert_similarity`
|
||||
raises `FileNotFoundError: Reference video folder does not exist`. The
|
||||
expected failure is what we want — the mp4 has already been written.
|
||||
- `--sync-generated-to-volume` + `--generated-volume-subdir`: copies the
|
||||
generated mp4s to the `hf-model-weights` Modal volume under
|
||||
`ssim_generated_videos/default/<SUBDIR>/generated_videos/` so we can pull
|
||||
them locally.
|
||||
|
||||
The Modal run will end with a nonzero exit (expected) and print a
|
||||
`modal volume get hf-model-weights ssim_generated_videos/default/<SUBDIR>/generated_videos ./generated_videos_modal/default`
|
||||
command. Capture that `<SUBDIR>` — you need it for step 3.
|
||||
|
||||
### 3. Download generated videos locally
|
||||
|
||||
```bash
|
||||
modal volume get --force hf-model-weights \
|
||||
ssim_generated_videos/default/"$SUBDIR"/generated_videos \
|
||||
./generated_videos_modal/default
|
||||
```
|
||||
|
||||
`--force` is required when the parent `./generated_videos_modal/default`
|
||||
already exists; without it, `modal volume get` errors with `[Errno 21] Is a
|
||||
directory`. Safe to pass on the first run too.
|
||||
|
||||
After this, the mp4s live at
|
||||
`./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
The extra `generated_videos/` level comes from the volume layout in
|
||||
`_sync_generated_videos_to_volume` (`ssim_test.py`) — the command copies
|
||||
`<repo>/fastvideo/tests/ssim/generated_videos/<tier>` to
|
||||
`ssim_generated_videos/<tier>/<SUBDIR>/generated_videos/`, and `modal volume
|
||||
get` preserves that trailing `generated_videos/` segment.
|
||||
|
||||
### 4. PAUSE — user reviews quality
|
||||
|
||||
Type-aware verification.
|
||||
|
||||
**For `ARTEFACT_TYPE = pixel`** — list the downloaded mp4s and ask the user to
|
||||
open them in a video player:
|
||||
|
||||
> "Generated videos downloaded to `./generated_videos_modal/default/generated_videos/L40S_reference_videos/`. Please open them and confirm the quality looks correct. Reply **`upload`** to continue, or anything else to abort."
|
||||
|
||||
**For `ARTEFACT_TYPE = latent`** — `.pt` files are not human-watchable. Print
|
||||
a numerics dump for each `.pt` so the user can sanity-check shape, distribution,
|
||||
and metadata:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from pathlib import Path
|
||||
ROOT = Path("./generated_videos_modal/default/generated_videos/L40S_reference_videos")
|
||||
for p in sorted(ROOT.rglob("*.pt")):
|
||||
d = torch.load(p, map_location="cpu", weights_only=False)
|
||||
s = d["expected_slice"]
|
||||
L = d["latent"].float()
|
||||
print(f"=== {p.relative_to(ROOT)} ===")
|
||||
print(f" format_version: {d['format_version']}")
|
||||
print(f" shape: {d['shape']}")
|
||||
print(f" dtype_original: {d['dtype_original']}")
|
||||
print(f" slice_spec: {d['slice_spec']}")
|
||||
print(f" slice shape={tuple(s.shape)} mean={s.mean():+.4f} std={s.std():.4f} min={s.min():+.4f} max={s.max():+.4f}")
|
||||
print(f" latent shape={tuple(L.shape)} mean={L.mean():+.4f} std={L.std():.4f} min={L.min():+.4f} max={L.max():+.4f}")
|
||||
print(f" finite: latent NaN={torch.isnan(L).any().item()} Inf={torch.isinf(L).any().item()}; "
|
||||
f"slice NaN={torch.isnan(s).any().item()} Inf={torch.isinf(s).any().item()}")
|
||||
print(f" metadata: {d['metadata']}\n")
|
||||
```
|
||||
|
||||
Sanity criteria:
|
||||
- `format_version == 1` (matches `LATENT_REFERENCE_FORMAT_VERSION`).
|
||||
- `shape` matches what the model produces (e.g. LTX-2 distilled =
|
||||
`[1, 128, T_lat, H_lat, W_lat]`; Stable Audio Open 1.0 = `[1, 64, 1024]`).
|
||||
- `slice_spec.kind` matches a registered kind (`corner_3x3_first_frame`
|
||||
for video, `audio_first_8_timesteps` for audio).
|
||||
- No `NaN`/`Inf`. `mean ≈ 0`, `std ≈ 1` (denoised latents stay close to
|
||||
the initial Gaussian distribution; very wide deviations suggest
|
||||
numerical drift).
|
||||
- `metadata.prompt` matches the test's prompt.
|
||||
|
||||
Then ask:
|
||||
|
||||
> "Numerics look right? Reply **`upload`** to continue, or anything else to abort."
|
||||
|
||||
Do not proceed until the user explicitly says `upload`. If they abort, leave
|
||||
everything on disk so they can inspect further — no cleanup.
|
||||
|
||||
### 5. Copy into the local reference layout
|
||||
|
||||
Scoped copy — only the new test's artefacts. Single command works for both
|
||||
artefact types because `_iter_reference_files` walks `.mp4` and `.pt`:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--generated-dir ./generated_videos_modal/default/generated_videos/L40S_reference_videos
|
||||
```
|
||||
|
||||
(The `--generated-dir` points at the device-folder root inside the
|
||||
downloaded tree; `copy-local` walks all `<model>/<backend>/*.{mp4,pt}`
|
||||
underneath it. Since the Modal run was scoped to a single test file via
|
||||
`--test-files`, only that test's model(s) are present — so the copy is
|
||||
implicitly per-test.)
|
||||
|
||||
Result for pixel: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
Result for latent: same path with `.pt` extension.
|
||||
|
||||
### 6. Upload to HF — scoped per model_id, with overwrite guard
|
||||
|
||||
For each `<model_id>`:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>"
|
||||
```
|
||||
|
||||
The upload command:
|
||||
|
||||
- Uploads **only** `reference_videos/default/L40S_reference_videos/<model_id>/`.
|
||||
- **Refuses** if any file already exists at that path on HF (this is the
|
||||
guard — seeding a new test should never clobber existing refs). To override,
|
||||
the user must re-run with `--force`. If the guard fires, stop and report
|
||||
exactly which files exist; do not silently `--force`.
|
||||
|
||||
Reads the HF token from `HF_API_KEY` / `HUGGINGFACE_HUB_TOKEN` / `HF_TOKEN`.
|
||||
|
||||
### 7. Report success
|
||||
|
||||
List what was uploaded (paths in repo) and remind the user to push any
|
||||
related code changes. Do **not** auto-verify by re-running Modal — the user
|
||||
can run `pytest fastvideo/tests/ssim/<test_file>` later to confirm end-to-end;
|
||||
it will auto-download the refs they just uploaded.
|
||||
|
||||
## Failure modes and how to handle them
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before step 2. The Modal run needs it (passed
|
||||
via `--hf-api-key`), and step 6 needs it for upload. If the user
|
||||
ran `hf auth login` instead of exporting an env var, read the cached
|
||||
token via `huggingface_hub.get_token()` and forward it to Modal as
|
||||
`--hf-api-key="$CACHED_TOKEN"`.
|
||||
- **Modal run fails before generation.** No artefacts on the volume — nothing
|
||||
to download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
|
||||
and retry from step 2.
|
||||
- **`./generated_videos_modal/default/L40S_reference_videos/` missing after
|
||||
`modal volume get`.** The run didn't produce artefacts (most likely the
|
||||
test crashed before writing, or `REQUIRED_GPUS` exceeded the partition
|
||||
capacity — see Modal logs).
|
||||
- **Latent test crashed with FSDP / inference_mode error
|
||||
(`RuntimeError: Inference tensors do not track version counter`).** The
|
||||
test must pass `init_kwargs_override={"use_fsdp_inference": False}` when
|
||||
`sp_size == 1` — see `test_stable_audio_similarity.py` for the pattern.
|
||||
Fix in the test, push, retry.
|
||||
- **Upload guard fires (files already exist).** The test name / model id
|
||||
collides with something already on HF. Verify the user actually wants to
|
||||
replace existing refs; if so, re-run the upload with `--force`. If not,
|
||||
rename the model id in `*_MODEL_TO_PARAMS` and re-seed.
|
||||
- **Quality looks wrong in step 4.** Abort. The artefacts stay on disk for
|
||||
inspection. The fix is usually in the test's params (resolution, steps,
|
||||
seed) — edit the test, then re-run the skill.
|
||||
- For latent: also check `slice_spec.kind` matches the latent rank
|
||||
(`corner_3x3_first_frame` requires 5-D, `audio_first_8_timesteps`
|
||||
requires 3-D); a rank/kind mismatch raises in `_extract_expected_slice`.
|
||||
|
||||
## Design notes (for future skill maintainers)
|
||||
|
||||
- The skill deliberately runs on Modal, **not** locally, because the CI
|
||||
runner is L40S. Seeding from a different GPU SKU produces refs that CI's
|
||||
L40S runs can't match (pixel SSIM drifts across SKUs; latent cosine has
|
||||
tighter cross-SKU bf16 drift but the configured tolerances assume
|
||||
same-SKU seed → same-SKU verify).
|
||||
- The skill is default-tier only. `full_quality` refs are seeded by a
|
||||
separate, deliberate operation — they double runtime and aren't what CI
|
||||
gates on.
|
||||
- The overwrite guard in `reference_videos_cli.py upload` is default-on
|
||||
specifically because this skill exists. Re-seeding is a distinct operation
|
||||
that requires explicit `--force`.
|
||||
- Both artefact types share the same Modal flow: the orchestrator sets
|
||||
`--skip-reference-download` + `--no-fail-fast`, runs pytest, the test's
|
||||
helper writes the artefact (`.mp4` via `imageio` for pixel,
|
||||
`save_latent_reference` → `torch.save` for latent) BEFORE the
|
||||
missing-reference assertion raises. `_sync_generated_videos_to_volume` in
|
||||
`ssim_test.py` does a `shutil.copytree` of the whole `generated_videos/`
|
||||
tree, picking up `.mp4`, `.pt`, and the `*_ssim.json` / `*_latent.json`
|
||||
metric files alongside.
|
||||
|
||||
## References
|
||||
|
||||
- `fastvideo/tests/modal/ssim_test.py` — Modal orchestrator; see
|
||||
`--sync-generated-to-volume`, `--generated-volume-subdir`,
|
||||
`--skip-reference-download`, `--no-fail-fast`.
|
||||
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
|
||||
(with `--model-id`, `--force`), `download`, `ensure` subcommands.
|
||||
Extension allowlist is `REFERENCE_EXTENSIONS = VIDEO_EXTENSIONS +
|
||||
LATENT_EXTENSIONS` (`.pt`).
|
||||
- `fastvideo/tests/ssim/README.md` — reference layout, HF repo conventions.
|
||||
- `fastvideo/tests/ssim/inference_similarity_utils.py` — pixel helpers
|
||||
(`run_text_to_video_similarity_test`,
|
||||
`run_image_to_video_similarity_test`, `build_init_kwargs`).
|
||||
- `fastvideo/tests/ssim/latent_similarity_utils.py` — latent helper
|
||||
(`run_text_to_latent_similarity_test`), slice spec dispatch
|
||||
(`_extract_expected_slice`), reference schema
|
||||
(`save_latent_reference` / `load_latent_reference`),
|
||||
`LATENT_REFERENCE_FORMAT_VERSION`.
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|------|--------|
|
||||
| 2026-04-17 | Initial version (Modal sync-to-volume flow). |
|
||||
| 2026-04-21 | Rewrite: single-test scope, explicit user-review pause, per-`model_id` upload, HF overwrite guard. Dropped `scripts/seed_ssim.sh`. |
|
||||
| 2026-04-21 | Post-first-run fixes: `modal volume get` needs `--force` when parent exists; download tree has an extra `generated_videos/` level so `--generated-dir` must reflect it. |
|
||||
| 2026-05-01 | Latent (`*.pt`) artefact support: artefact-type detection in step 1, type-aware verification (visual eyeball for mp4, numerics dump for pt) in step 4, FSDP+inference_mode failure-mode added, design notes for the unified Modal flow. Triggered by PR #1253 (LTX-2 latent migration + Stable Audio latent test). |
|
||||
@@ -29,8 +29,8 @@
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
],
|
||||
"run_config": {
|
||||
"num_warmup_runs": 1,
|
||||
"num_measurement_runs": 3,
|
||||
"num_warmup_runs": 2,
|
||||
"num_measurement_runs": 5,
|
||||
"required_gpus": 2
|
||||
},
|
||||
"thresholds": {
|
||||
|
||||
@@ -15,8 +15,21 @@ log "Project root: $PROJECT_ROOT"
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
python3 -m pip install modal
|
||||
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "uv not found, bootstrapping..."
|
||||
if ! curl -LsSf https://astral.sh/uv/install.sh | sh; then
|
||||
log "Error: Failed to bootstrap uv via astral.sh installer."
|
||||
exit 1
|
||||
fi
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "Error: uv still not on PATH after bootstrap."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
# --break-system-packages preserves prior `pip install --user` semantics on PEP 668 agents.
|
||||
uv pip install --system --break-system-packages modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
@@ -63,7 +76,72 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR IMAGE_VERSION=$IMAGE_VERSION"
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
upload_performance_artifacts() {
|
||||
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
|
||||
LOCAL_DIR="downloaded_reports"
|
||||
|
||||
_download_reports() {
|
||||
log "Downloading perf_reports/ from Modal Volume..."
|
||||
mkdir -p "$LOCAL_DIR"
|
||||
if ! modal volume get hf-model-weights "perf_reports/" "$LOCAL_DIR"; then
|
||||
log "Error: Failed to download perf_reports/ from Modal Volume."
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_dashboard() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_${SHORT_SHA}_*" | head -n 1)
|
||||
log "TARGET dashboard: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
log "Found dashboard: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
buildkite-agent annotate --style info --context "perf-dashboard" < "$target"
|
||||
else
|
||||
log "Warning: Could not find a dashboard file matching $SHORT_SHA"
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_perf_summary() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "perf_${SHORT_SHA}_*" | head -n 1)
|
||||
log "TARGET perf summary: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
log "Found perf summary: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
buildkite-agent annotate --style info --context "perf-summary" < "$target"
|
||||
else
|
||||
log "Warning: Could not find a perf summary file matching $SHORT_SHA"
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_modal_volume() {
|
||||
log "Cleaning up perf_reports/ from Modal Volume..."
|
||||
if modal volume rm hf-model-weights "perf_reports/" --recursive; then
|
||||
log "Successfully deleted perf_reports/ from Modal Volume."
|
||||
else
|
||||
log "Warning: Failed to delete perf_reports/ from Modal Volume. Manual cleanup may be required."
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_local() {
|
||||
log "Cleaning up local download directory..."
|
||||
rm -rf "$LOCAL_DIR"
|
||||
}
|
||||
|
||||
# --- Main flow ---
|
||||
_download_reports || { _cleanup_local; return 1; }
|
||||
_upload_dashboard
|
||||
_upload_perf_summary
|
||||
_cleanup_modal_volume
|
||||
_cleanup_local
|
||||
}
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
@@ -124,8 +202,9 @@ case "$TEST_TYPE" in
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
|
||||
;;
|
||||
"performance")
|
||||
log "Running performance tests..."
|
||||
log "Running performance tests on Modal..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_performance_tests"
|
||||
POST_RUN_HOOK="upload_performance_artifacts"
|
||||
;;
|
||||
"api_server")
|
||||
log "Running API server integration tests..."
|
||||
@@ -147,5 +226,10 @@ else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
if [ -n "$POST_RUN_HOOK" ]; then
|
||||
log "Executing post-run hook: $POST_RUN_HOOK"
|
||||
"$POST_RUN_HOOK"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
|
||||
@@ -13,8 +13,21 @@ log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "uv not found, bootstrapping..."
|
||||
if ! curl -LsSf https://astral.sh/uv/install.sh | sh; then
|
||||
log "Error: Failed to bootstrap uv via astral.sh installer."
|
||||
exit 1
|
||||
fi
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "Error: uv still not on PATH after bootstrap."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
# --break-system-packages preserves prior `pip install --user` semantics on PEP 668 agents.
|
||||
uv pip install --system --break-system-packages pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
|
||||
+1
-1
@@ -105,7 +105,7 @@ pull_request_rules:
|
||||
- files~=^fastvideo/pipelines/samplers/
|
||||
- files~=^fastvideo/entrypoints/
|
||||
- files~=^fastvideo/worker/
|
||||
- files~=^fastvideo/configs/sample/
|
||||
- files~=^fastvideo/api/sampling_param
|
||||
- files~=^fastvideo/configs/pipelines/
|
||||
- files~=^examples/inference/
|
||||
- -closed
|
||||
|
||||
@@ -37,10 +37,11 @@ jobs:
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-mkdocs.txt
|
||||
run: uv pip install --system -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
@@ -56,10 +56,11 @@ jobs:
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install build dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install build twine wheel
|
||||
run: uv pip install --system build twine wheel
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
|
||||
@@ -131,11 +131,13 @@ jobs:
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
uv pip install --system typing-extensions==4.12.2
|
||||
uv pip install --system --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
@@ -145,20 +147,20 @@ jobs:
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
|
||||
uv pip install --system setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
cd fastvideo-kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
# Release builds are produced on GPU-less runners, so force-enable TK and target Hopper.
|
||||
export TORCH_CUDA_ARCH_LIST="9.0a"
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
|
||||
|
||||
|
||||
# Build standard wheel (no local version suffix) for PyPI
|
||||
python -m build --wheel --outdir dist
|
||||
|
||||
|
||||
# Fix the wheel to be manylinux compliant
|
||||
pip install auditwheel
|
||||
uv pip install --system auditwheel
|
||||
# Point auditwheel at torch libs, but do not vendor them into the wheel.
|
||||
TORCH_LIB_DIR=$(python - <<'PY'
|
||||
import os
|
||||
@@ -211,10 +213,13 @@ jobs:
|
||||
pattern: 'fastvideo_kernel-py*'
|
||||
merge-multiple: true
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
pip install build scikit-build-core cmake ninja
|
||||
|
||||
uv pip install --system build scikit-build-core cmake ninja
|
||||
|
||||
cd fastvideo-kernel
|
||||
# We don't need full CUDA/Torch to just package the source (sdist)
|
||||
python -m build --sdist --outdir dist
|
||||
|
||||
@@ -92,3 +92,11 @@ preprocess_output_text/
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
|
||||
# Local clones of upstream repos used only for parity testing.
|
||||
/stable-audio-tools/
|
||||
/daVinci-MagiHuman/
|
||||
|
||||
# Converted model weights (produced by scripts/checkpoint_conversion/*).
|
||||
# Tens of GB; should live on HF, not in git.
|
||||
/converted_weights/
|
||||
|
||||
@@ -7,20 +7,13 @@ exclude: |
|
||||
fastvideo-kernel/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
assets/prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
examples/.*|
|
||||
\.agents/.*|
|
||||
.github/workflows/publish-fastvideo.yml|
|
||||
.github/workflows/_template-build-image.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
.github/workflows/_template-build-image.yml
|
||||
)
|
||||
repos:
|
||||
- repo: https://github.com/google/yapf
|
||||
|
||||
@@ -11,7 +11,7 @@
|
||||
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
- `uv pip install -e .[dev]`: editable install with lint/test extras.
|
||||
- `uv pip install -e ".[dev]"`: editable install with lint/test extras.
|
||||
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
|
||||
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
|
||||
- `pytest tests/`: run top-level test suite.
|
||||
@@ -23,7 +23,8 @@
|
||||
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
|
||||
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
|
||||
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
|
||||
- Target line length is 80.
|
||||
- Lint via `pre-commit run --files <changed paths>` (or `pre-commit run --all-files` for a full sweep) before committing. Do not shell out to `yapf`/`ruff`/`codespell`/`mypy` directly — pre-commit chains them with the project's config and respects the `.pre-commit-config.yaml` excludes (e.g. `fastvideo/tests/` is intentionally skipped). If pre-commit reports `(no files to check)` for your paths, that exclude is deliberate — don't bypass it.
|
||||
- Target line length is 120 (configured in `pyproject.toml` for ruff, yapf, and isort).
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
@@ -54,3 +55,31 @@ This repository is agent-friendly. Before doing any work, read:
|
||||
If you are exploring a new procedure that has no existing SOP, document your
|
||||
progress in `.agents/exploration/` and flag it for review at the end of your
|
||||
session.
|
||||
|
||||
## Per-Directory AGENTS.md
|
||||
|
||||
Local guidance lives next to the code. Read the in-scope file before editing:
|
||||
|
||||
| Directory | What it covers |
|
||||
|-----------|----------------|
|
||||
| `fastvideo/AGENTS.md` | Core package map, public API, registry-driven model dispatch |
|
||||
| `fastvideo/configs/AGENTS.md` | Arch + pipeline config dataclasses, `param_names_mapping` |
|
||||
| `fastvideo/models/AGENTS.md` | DiT / VAE / encoder / scheduler / loader layout (pre-commit excluded) |
|
||||
| `fastvideo/layers/AGENTS.md` | Tensor-parallel linear/attention layer rules for ports |
|
||||
| `fastvideo/attention/AGENTS.md` | Backend registry + env-var override |
|
||||
| `fastvideo/pipelines/AGENTS.md` | Stage ABC, `basic/<model>/`, `preprocess/`, presets |
|
||||
| `fastvideo/training/AGENTS.md` | Legacy monolithic pipelines (frozen for existing models) |
|
||||
| `fastvideo/train/AGENTS.md` | New modular trainer (methods × models × callbacks, YAML) |
|
||||
| `fastvideo/tests/AGENTS.md` | Test taxonomy, conftest, pre-commit-excluded path |
|
||||
| `fastvideo/tests/ssim/AGENTS.md` | GPU SSIM regression authoring + reference video sync |
|
||||
| `scripts/checkpoint_conversion/AGENTS.md` | Adding a converter for a new HF/official checkpoint |
|
||||
|
||||
## Critical: Two Training Stacks Coexist
|
||||
|
||||
- `fastvideo/training/` — legacy, monolithic per-model `*_training_pipeline.py` and
|
||||
`*_distillation_pipeline.py`. Still authoritative for shipped models.
|
||||
- `fastvideo/train/` — new modular framework (composable methods × models × callbacks
|
||||
driven by YAML). Preferred for new training work.
|
||||
|
||||
Pick the matching stack before editing. Do not migrate a pipeline between them
|
||||
without an explicit ask — the conventions and config surfaces differ.
|
||||
|
||||
@@ -128,7 +128,7 @@ class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "openai/clip-vit-base-patch32"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
raise ImportError("Please install transformers: uv pip install transformers")
|
||||
super().__init__(device)
|
||||
self.processor = CLIPProcessor.from_pretrained(model_name)
|
||||
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
|
||||
@@ -171,7 +171,7 @@ class VideoMAEFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "MCG-NJU/videomae-base"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
raise ImportError("Please install transformers: uv pip install transformers")
|
||||
super().__init__(device)
|
||||
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
@@ -57,7 +57,7 @@ class I3DFeatureExtractor(nn.Module):
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
f"uv pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless transformers huggingface_hub
|
||||
uv pip install -q opencv-python-headless transformers huggingface_hub
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
uv pip install -q opencv-python-headless
|
||||
+2
-2
@@ -38,10 +38,10 @@ cp -r /path/to/FastVideo/comfyui /path/to/ComfyUI/custom_nodes/FastVideo
|
||||
|
||||
#### Install dependencies:
|
||||
|
||||
Currently, the only dependency is `fastvideo`, which can be installed using pip.
|
||||
Currently, the only dependency is `fastvideo`, which can be installed with `uv`.
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
#### Install missing custom nodes:
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp310-cp310-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp310-cp310-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.11 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp311-cp311-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp311-cp311-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp312-cp312-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,7 +42,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
@@ -50,7 +50,7 @@ COPY . .
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -43,7 +43,7 @@ COPY . .
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[rocm] && \
|
||||
uv pip install --no-cache-dir -e ".[rocm]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r requirements-mkdocs.txt
|
||||
uv pip install -r requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -62,9 +62,9 @@ This page contains the complete API reference for the FastVideo library.
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.configs.sample
|
||||
#### fastvideo.api.sampling_param
|
||||
|
||||
::: fastvideo.configs.sample
|
||||
::: fastvideo.api.sampling_param
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
|
||||
@@ -0,0 +1,210 @@
|
||||
# Activation Trace Mode
|
||||
|
||||
!!! note
|
||||
This page covers Extension 0 (module forward hooks), which is the implemented
|
||||
tracing mechanism. Extensions 1-3 are design sketches for future work and are
|
||||
**not yet implemented**.
|
||||
|
||||
## Overview
|
||||
|
||||
Activation trace mode is a zero-overhead-when-off, env-gated mechanism for
|
||||
dumping per-layer activation statistics during FastVideo inference. Its primary
|
||||
use case is **parity debugging across model ports**: enable tracing on both
|
||||
FastVideo and the upstream reference implementation, then `diff` the resulting
|
||||
JSONL files to find the first divergent layer.
|
||||
|
||||
The mechanism is intentionally narrow. It doesn't replace general logging,
|
||||
profiling, or function tracing. It answers one question: "at which layer do
|
||||
FastVideo and the reference model first produce different numbers?"
|
||||
|
||||
## When to use
|
||||
|
||||
- Investigating numerical drift between FastVideo and an upstream reference.
|
||||
- Debugging mid-pipeline divergence (e.g., one block produces wrong output while earlier blocks match).
|
||||
- Validating that a refactor preserves bf16 noise-floor behavior across many layers.
|
||||
|
||||
## When NOT to use
|
||||
|
||||
| Goal | Use instead |
|
||||
|---|---|
|
||||
| General logging | `init_logger(__name__)` |
|
||||
| Per-stage timing | `FASTVIDEO_STAGE_LOGGING` |
|
||||
| Profiling kernel timings | `FASTVIDEO_TORCH_PROFILER_DIR` (see [Profiling](profiling.md)) |
|
||||
| Function-call tracing | `FASTVIDEO_TRACE_FUNCTION` (heavy) |
|
||||
|
||||
## Quickstart
|
||||
|
||||
```bash
|
||||
FASTVIDEO_TRACE_ACTIVATIONS=1 \
|
||||
FASTVIDEO_TRACE_LAYERS="^block\.layers\.[0-9]+$" \
|
||||
FASTVIDEO_TRACE_STATS="abs_mean,sum,max,shape" \
|
||||
FASTVIDEO_TRACE_OUTPUT="/tmp/fv_trace.jsonl" \
|
||||
python examples/inference/basic/basic_magi_human.py
|
||||
```
|
||||
|
||||
Each line in `/tmp/fv_trace.jsonl` is a JSON record:
|
||||
|
||||
```json
|
||||
{"module": "block.layers.0", "tensor": "out", "step": 0, "abs_mean": 1.234, "sum": -5.678, "max": 9.012, "shape": [1, 4096, 5120]}
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
| Env var | Default | Description |
|
||||
|---|---|---|
|
||||
| `FASTVIDEO_TRACE_ACTIVATIONS` | `False` | Master toggle. When unset or false, **zero overhead** in the production hot path. |
|
||||
| `FASTVIDEO_TRACE_LAYERS` | `""` (all) | Python regex filter applied to `model.named_modules()` names. Empty string matches all modules. |
|
||||
| `FASTVIDEO_TRACE_STATS` | `"abs_mean,sum"` | Comma-separated stats to compute. Available: `abs_mean`, `sum`, `min`, `max`, `mean`, `std`, `shape`, `dtype`. |
|
||||
| `FASTVIDEO_TRACE_OUTPUT` | `"/tmp/fv_trace_<pid>.jsonl"` | Output file path. `<pid>` is replaced with the process ID at runtime. |
|
||||
| `FASTVIDEO_TRACE_STEPS` | `""` (all) | Comma-separated denoising step indices to capture. Empty string captures all steps. |
|
||||
|
||||
## Workflow: parity-debug a model port
|
||||
|
||||
1. Set up a tightly-controlled comparison: a parity test or a small standalone
|
||||
script that loads both the FastVideo model and the upstream reference with
|
||||
identical inputs and seeds.
|
||||
|
||||
2. Run the FastVideo side with tracing on:
|
||||
|
||||
```bash
|
||||
FASTVIDEO_TRACE_ACTIVATIONS=1 \
|
||||
FASTVIDEO_TRACE_LAYERS="<your regex>" \
|
||||
FASTVIDEO_TRACE_OUTPUT="/tmp/fv_trace_fv.jsonl" \
|
||||
python <fv_runner.py>
|
||||
```
|
||||
|
||||
3. Run the upstream side. The upstream repo needs separate instrumentation. See
|
||||
"Hooking the upstream side" below.
|
||||
|
||||
4. Sort both files by `(module, step)` if needed, then diff:
|
||||
|
||||
```bash
|
||||
diff /tmp/fv_trace_fv.jsonl /tmp/fv_trace_upstream.jsonl
|
||||
```
|
||||
|
||||
5. The first divergent line identifies the first layer where FastVideo and the
|
||||
upstream produce different outputs. Start debugging there.
|
||||
|
||||
## Architecture (Extension 0: module forward hooks)
|
||||
|
||||
At pipeline initialization, `attach_activation_trace()` reads the env vars once.
|
||||
If `FASTVIDEO_TRACE_ACTIVATIONS` is unset or false, the function returns
|
||||
immediately and no hooks are registered. If tracing is on, it walks
|
||||
`model.named_modules()`, filters by the layer regex, and registers an
|
||||
`ActivationStatHook` on each matching module.
|
||||
|
||||
During the forward pass, each hook fires after its module completes, computes
|
||||
the requested stats on the output tensor, and appends a JSON record to the
|
||||
output file.
|
||||
|
||||
```
|
||||
ComposedPipelineBase
|
||||
└─ attach_activation_trace()
|
||||
├─ reads env vars (once at startup)
|
||||
├─ if off: returns None immediately
|
||||
└─ if on: walks named_modules()
|
||||
└─ registers ActivationStatHook on matching modules
|
||||
└─ on each forward: compute stats → append JSONL
|
||||
```
|
||||
|
||||
### Zero-overhead-when-off guarantee
|
||||
|
||||
- The env var check happens **once at startup** inside `attach_activation_trace()`.
|
||||
- If the env var is unset or false, the function returns `None` immediately.
|
||||
- No hooks are registered. No branches are added to the production forward path.
|
||||
- The only cost when tracing is off is one env var lookup at pipeline
|
||||
initialization, which takes under a microsecond.
|
||||
|
||||
### Hooking the upstream side
|
||||
|
||||
The upstream reference repo isn't part of FastVideo, so it can't read FastVideo
|
||||
env vars directly. Two options:
|
||||
|
||||
**Option 1: Inline patch** in your local clone of the upstream repo. Add
|
||||
`register_forward_hook` calls in the same shape as `ActivationStatHook`. Clean
|
||||
up afterward with `git stash` or `git checkout HEAD -- <file>`.
|
||||
|
||||
**Option 2: Wrapper script**. Write a small Python harness that imports the
|
||||
upstream model, walks its `named_modules()`, and attaches hooks externally.
|
||||
This is the same pattern used in
|
||||
`tests/local_tests/transformers/_debug_magi_human_block_parity.py`.
|
||||
|
||||
The `add-model-trace` skill at `~/.config/opencode/skill/add-model-trace/`
|
||||
provides a script template for this purpose.
|
||||
|
||||
## Future extensions (design only, not yet implemented)
|
||||
|
||||
### Extension 1: FX/Dynamo backend graph rewrite
|
||||
|
||||
**Granularity**: per-FX-node (every matmul, every add).
|
||||
|
||||
**Mechanism**: a `torch.compile` backend that takes the captured `GraphModule`
|
||||
and inserts logger nodes after each op. Compiles into a separate artifact from
|
||||
the production graph.
|
||||
|
||||
**Off semantics**: zero overhead. The production compile path is untouched.
|
||||
|
||||
**When to add**: if you need to trace inside a `torch.compile`'d graph and
|
||||
Extension 0 is too coarse.
|
||||
|
||||
**Build cost**: roughly 1-2 days. Reference:
|
||||
`torchao.quantization.pt2e._numeric_debugger`.
|
||||
|
||||
### Extension 2: AST source injection at import time
|
||||
|
||||
**Granularity**: per-line (between any two Python statements).
|
||||
|
||||
**Mechanism**: an importlib loader hook rewrites Python source AST at module
|
||||
import time, inserting `if TRACE: dump(...)` statements. The decision is made
|
||||
once at import.
|
||||
|
||||
**Off semantics**: zero overhead. If the env var is off at import time, source
|
||||
is loaded as-is.
|
||||
|
||||
**When to add**: if you need per-line granularity that even FX-node-level can't
|
||||
provide. This is almost never the right choice.
|
||||
|
||||
**Build cost**: roughly 1 week. Brittle and hard to debug.
|
||||
|
||||
### Extension 3: `__torch_dispatch__` / `TorchDispatchMode`
|
||||
|
||||
**Granularity**: per-op (every dispatcher call: matmul, add, view, etc.).
|
||||
|
||||
**Mechanism**: a `TorchDispatchMode` context manager that intercepts all ops at
|
||||
the dispatcher level.
|
||||
|
||||
**Off semantics**: zero overhead. PyTorch's dispatcher only invokes mode hooks
|
||||
when a mode is active.
|
||||
|
||||
**When on**: significant overhead. Every op pays a Python callback cost. Triton
|
||||
kernels bypass it.
|
||||
|
||||
**When to add**: useful for quantization or dtype debugging where module-level
|
||||
granularity isn't enough.
|
||||
|
||||
**Build cost**: roughly 1 day. Reference:
|
||||
`torch.utils._python_dispatch.TorchDispatchMode`.
|
||||
|
||||
## Comparison with similar tools
|
||||
|
||||
| Tool | Pattern | FastVideo equivalent |
|
||||
|---|---|---|
|
||||
| SGLang `--debug-tensor-dump-output-folder` | env-gated forward hooks at startup | Extension 0 (this) |
|
||||
| TransformerEngine `DumpTensors` | config-driven selective dumps | Extension 0 (env-driven) |
|
||||
| HuggingFace `output_hidden_states=True` | source-level boolean gating | Not used; Extension 0 avoids model code edits |
|
||||
| torchao numeric debugger | FX pass + node-level loggers | Extension 1 (future) |
|
||||
| W&B `wandb.watch()` | runtime forward hooks (always on once registered) | Extension 0 has a similar mechanism, but gated off by default |
|
||||
|
||||
## Implementation references
|
||||
|
||||
- Module: `fastvideo/hooks/activation_trace.py`
|
||||
- Env vars: `fastvideo/envs.py` (`FASTVIDEO_TRACE_ACTIVATIONS` and friends)
|
||||
- Pipeline integration: `fastvideo/pipelines/composed_pipeline_base.py`
|
||||
- Tests: `fastvideo/tests/hooks/test_activation_trace.py`
|
||||
- Companion skill (for ad-hoc port investigations): `~/.config/opencode/skill/add-model-trace/`
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|---|---|
|
||||
| 2026-05-01 | Initial Extension 0 (module forward hooks) implementation. Extensions 1-3 designed but not implemented. |
|
||||
@@ -173,7 +173,7 @@ Applied by Mergify based on which paths you modified. Multiple scope labels can
|
||||
| Label | File paths that trigger it |
|
||||
|-------|---------------------------|
|
||||
| `scope: training` | `fastvideo/train/`, `fastvideo/training/`, `fastvideo/distillation/`, `examples/train/`, `examples/training/`, `examples/distill/` |
|
||||
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/configs/sample/`, `fastvideo/configs/pipelines/`, `examples/inference/` |
|
||||
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/api/sampling_param.py`, `fastvideo/configs/pipelines/`, `examples/inference/` |
|
||||
| `scope: attention` | `fastvideo/attention/` |
|
||||
| `scope: kernel` | `fastvideo-kernel/`, `csrc/` |
|
||||
| `scope: data` | `fastvideo/dataset/`, `fastvideo/pipelines/preprocess/`, `examples/preprocessing/` |
|
||||
|
||||
@@ -44,7 +44,7 @@ FastVideo maps a Diffusers-style repo into a pipeline like:
|
||||
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
|
||||
weight name translation.
|
||||
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
|
||||
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
|
||||
- `fastvideo/api/sampling_param.py`: runtime sampling parameters.
|
||||
- `fastvideo/pipelines/basic/*`: end-to-end pipeline logic built from stages.
|
||||
- `model_index.json`: the HF repo entrypoint that maps component names to
|
||||
classes and weight files.
|
||||
@@ -55,7 +55,7 @@ Minimal usage example (based on `examples/inference/basic/basic.py`):
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
|
||||
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
|
||||
@@ -296,8 +296,10 @@ Action:
|
||||
|
||||
- Add or reuse a numerical parity test that loads the official model and the
|
||||
FastVideo model and compares outputs.
|
||||
- See examples in `tests/local_tests/` (e.g., `tests/local_tests/upsamplers/`)
|
||||
and the commands in `tests/local_tests/README.md`.
|
||||
- See examples in `tests/local_tests/` organized by model family
|
||||
(e.g., `tests/local_tests/sd35/`, `tests/local_tests/ltx2/`,
|
||||
`tests/local_tests/stable_audio/`) and the navigation index in
|
||||
`tests/local_tests/README.md`.
|
||||
- If there are discrepancies, add opt‑in logging to both models and compare
|
||||
activation summaries (layer output sums, per‑stage logs).
|
||||
- First align the loaded weights (validate `param_names_mapping`).
|
||||
@@ -319,7 +321,8 @@ Purpose:
|
||||
|
||||
- `fastvideo/configs/pipelines/` describes pipeline wiring and model module
|
||||
names.
|
||||
- `fastvideo/configs/sample/` defines default runtime parameters.
|
||||
- `fastvideo/api/sampling_param.py` defines runtime sampling parameters.
|
||||
Defaults come from profiles in `fastvideo/pipelines/basic/<family>/profiles.py`.
|
||||
|
||||
Action:
|
||||
|
||||
@@ -347,7 +350,8 @@ Purpose:
|
||||
|
||||
Action:
|
||||
|
||||
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
|
||||
- Add a pipeline parity test under `tests/local_tests/<family>/`
|
||||
(e.g., `tests/local_tests/<family>/test_<family>_pipeline_parity.py`).
|
||||
- See the [Testing Guide](testing.md) for test conventions.
|
||||
|
||||
### 7) Add user‑facing examples
|
||||
@@ -474,7 +478,7 @@ FastVideo integration.
|
||||
3. Pipeline wiring.
|
||||
- Pipeline: `fastvideo/pipelines/basic/wan/wan_pipeline.py`
|
||||
- Pipeline config: `fastvideo/configs/pipelines/wan.py`
|
||||
- Sampling defaults: `fastvideo/configs/sample/wan.py`
|
||||
- Sampling defaults: `fastvideo/pipelines/basic/wan/profiles.py`
|
||||
|
||||
4. Minimal example.
|
||||
- Script: `examples/inference/basic/basic.py`
|
||||
|
||||
@@ -99,7 +99,7 @@ cd /FastVideo
|
||||
**Install the package**
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
uv pip install -e ".[dev]"
|
||||
```
|
||||
|
||||
The Docker image already includes Flash Attention and most heavy dependencies, so this is fast.
|
||||
|
||||
@@ -49,7 +49,7 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Install FastVideo in editable mode and set up hooks:
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
uv pip install -e ".[dev]"
|
||||
|
||||
# Optional: FlashAttention (builds native kernels)
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
status_definitions:
|
||||
kept: "Public field remains on a public adapter surface with the same meaning."
|
||||
moved: "Public field remains supported but normalizes into a different nested path."
|
||||
profile_owned: "Public field remains supported only through a model/profile-specific surface."
|
||||
preset_owned: "Public field remains supported only through a model/preset-specific surface."
|
||||
compatibility_only: "Legacy public field remains adapter-only during migration and is not part of the canonical typed schema."
|
||||
private_only: "Field should only be handled by private adapters and is not a public FastVideo compatibility promise."
|
||||
internal_only: "Field is runtime/config plumbing and should not be part of the new public typed inference API."
|
||||
@@ -29,7 +29,7 @@ surfaces:
|
||||
vae_cpu_offload: generator.engine.offload.vae
|
||||
pin_cpu_memory: generator.engine.offload.pin_cpu_memory
|
||||
enable_torch_compile: generator.engine.compile.enabled
|
||||
torch_compile_kwargs: generator.engine.compile.kwargs
|
||||
torch_compile_kwargs: generator.engine.compile.backend,fullgraph,mode,dynamic,extras
|
||||
disable_autocast: generator.engine.disable_autocast
|
||||
enable_stage_verification: generator.engine.enable_stage_verification
|
||||
prompt_txt: request.inputs.prompt_path
|
||||
@@ -40,12 +40,12 @@ surfaces:
|
||||
init_weights_from_safetensors_2: generator.pipeline.components.transformer_2_weights
|
||||
override_pipeline_cls_name: generator.pipeline.components.override_pipeline_cls_name
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
profile_owned:
|
||||
ltx2_vae_tiling: generator.pipeline.profile_overrides.ltx2.vae_tiling
|
||||
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_size_in_pixels
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
|
||||
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_size_in_frames
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_overlap_in_frames
|
||||
ltx2_vae_tiling: generator.pipeline.vae_tiling
|
||||
preset_owned:
|
||||
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.preset_overrides.ltx2.vae.spatial_tile_size_in_pixels
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.preset_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
|
||||
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.preset_overrides.ltx2.vae.temporal_tile_size_in_frames
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.preset_overrides.ltx2.vae.temporal_tile_overlap_in_frames
|
||||
ltx2_initial_latent_path: request.extensions.ltx2.initial_latent_path
|
||||
compatibility_only:
|
||||
mode: "Legacy multi-mode FastVideoArgs switch; typed inference config should not expose execution mode."
|
||||
@@ -69,16 +69,16 @@ surfaces:
|
||||
pipeline_config_base:
|
||||
moved:
|
||||
pipeline_config_path: generator.pipeline.components.pipeline_config_path
|
||||
profile_owned:
|
||||
embedded_cfg_scale: generator.pipeline.profile_overrides.embedded_cfg_scale
|
||||
flow_shift: generator.pipeline.profile_overrides.flow_shift
|
||||
flow_shift_sr: generator.pipeline.profile_overrides.flow_shift_sr
|
||||
is_causal: generator.pipeline.profile_overrides.is_causal
|
||||
vae_tiling: generator.pipeline.profile_overrides.vae_tiling
|
||||
vae_sp: generator.pipeline.profile_overrides.vae_sp
|
||||
dmd_denoising_steps: generator.pipeline.profile_overrides.dmd_denoising_steps
|
||||
ti2v_task: generator.pipeline.profile_overrides.ti2v_task
|
||||
boundary_ratio: generator.pipeline.profile_overrides.boundary_ratio
|
||||
preset_owned:
|
||||
embedded_cfg_scale: generator.pipeline.preset_overrides.embedded_cfg_scale
|
||||
flow_shift: generator.pipeline.preset_overrides.flow_shift
|
||||
flow_shift_sr: generator.pipeline.preset_overrides.flow_shift_sr
|
||||
is_causal: generator.pipeline.preset_overrides.is_causal
|
||||
vae_tiling: generator.pipeline.preset_overrides.vae_tiling
|
||||
vae_sp: generator.pipeline.preset_overrides.vae_sp
|
||||
dmd_denoising_steps: generator.pipeline.preset_overrides.dmd_denoising_steps
|
||||
ti2v_task: generator.pipeline.preset_overrides.ti2v_task
|
||||
boundary_ratio: generator.pipeline.preset_overrides.boundary_ratio
|
||||
compatibility_only:
|
||||
model_path: "Redundant with generator.model_path."
|
||||
disable_autocast: "Duplicated by generator.engine.disable_autocast during migration."
|
||||
@@ -97,7 +97,7 @@ surfaces:
|
||||
postprocess_text_funcs: "Internal text postprocessing hooks."
|
||||
|
||||
pipeline_config_extensions:
|
||||
profile_owned:
|
||||
preset_owned:
|
||||
conditioning_strategy:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
@@ -306,11 +306,35 @@ surfaces:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
num_frames_per_block:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
audio_channels:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
audio_end_in_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
audio_start_in_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
max_audio_duration_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
sample_size:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
sampling_rate:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
compatibility_only:
|
||||
batch_size: "Gen3C inference-only tuning field pending typed batching design."
|
||||
gradient_checkpointing: "Gen3C inference-only compatibility field pending typed batching design."
|
||||
guidance_scale: "Gen3C pipeline-level default pending profile/default-request cleanup."
|
||||
num_inference_steps: "Gen3C pipeline-level default pending profile/default-request cleanup."
|
||||
guidance_scale: "Gen3C pipeline-level default pending preset/default-request cleanup."
|
||||
num_inference_steps: "Gen3C pipeline-level default pending preset/default-request cleanup."
|
||||
internal_only:
|
||||
audio_decoder_config: "Legacy internal component config object."
|
||||
audio_decoder_precision: "Precision override pending dedicated component precision design."
|
||||
@@ -345,6 +369,7 @@ surfaces:
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
num_inference_steps_sr: request.sampling.num_inference_steps_sr
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
@@ -353,96 +378,43 @@ surfaces:
|
||||
return_frames: request.output.return_frames
|
||||
return_trajectory_latents: request.runtime.return_trajectory_latents
|
||||
return_trajectory_decoded: request.runtime.return_trajectory_decoded
|
||||
profile_owned:
|
||||
continuation_state: request.state
|
||||
return_continuation_state: request.output.return_state
|
||||
preset_owned:
|
||||
t_thresh: request.stage_overrides.refine.t_thresh
|
||||
spatial_refine_only: request.stage_overrides.refine.spatial_refine_only
|
||||
num_cond_frames: request.stage_overrides.refine.num_cond_frames
|
||||
trajectory_type: request.extensions.gen3c.trajectory_type
|
||||
movement_distance: request.extensions.gen3c.movement_distance
|
||||
camera_rotation: request.extensions.gen3c.camera_rotation
|
||||
prompt_attention_mask: request.extensions.hyworld.prompt_attention_mask
|
||||
negative_attention_mask: request.extensions.hyworld.negative_attention_mask
|
||||
camera_states: request.extensions.hunyuangamecraft.camera_states
|
||||
camera_trajectory: request.extensions.hunyuangamecraft.camera_trajectory
|
||||
action_list: request.extensions.hunyuangamecraft.action_list
|
||||
action_speed_list: request.extensions.hunyuangamecraft.action_speed_list
|
||||
gt_latents: request.extensions.hunyuangamecraft.gt_latents
|
||||
conditioning_mask: request.extensions.hunyuangamecraft.conditioning_mask
|
||||
ltx2_cfg_scale_video: request.extensions.ltx2.cfg_scale_video
|
||||
ltx2_cfg_scale_audio: request.extensions.ltx2.cfg_scale_audio
|
||||
ltx2_modality_scale_video: request.extensions.ltx2.modality_scale_video
|
||||
ltx2_modality_scale_audio: request.extensions.ltx2.modality_scale_audio
|
||||
ltx2_rescale_scale: request.extensions.ltx2.rescale_scale
|
||||
ltx2_stg_scale_video: request.extensions.ltx2.stg_scale_video
|
||||
ltx2_stg_scale_audio: request.extensions.ltx2.stg_scale_audio
|
||||
ltx2_stg_blocks_video: request.extensions.ltx2.stg_blocks_video
|
||||
ltx2_stg_blocks_audio: request.extensions.ltx2.stg_blocks_audio
|
||||
audio_start_in_s: request.extensions.stable_audio.audio_start_in_s
|
||||
audio_end_in_s: request.extensions.stable_audio.audio_end_in_s
|
||||
init_audio: request.extensions.stable_audio.init_audio
|
||||
init_audio_strength: request.extensions.stable_audio.init_audio_strength
|
||||
init_noise_level: request.extensions.stable_audio.init_noise_level
|
||||
inpaint_audio: request.extensions.stable_audio.inpaint_audio
|
||||
inpaint_mask: request.extensions.stable_audio.inpaint_mask
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
|
||||
sampling_param_extensions:
|
||||
moved:
|
||||
guidance_scale_2:
|
||||
target: request.sampling.guidance_scale_2
|
||||
sources:
|
||||
- fastvideo.configs.sample.lingbotworld.LingBotWorld_SamplingParam
|
||||
- fastvideo.configs.sample.lingbotworld.Wan2_2_I2V_A14B_SamplingParam
|
||||
- fastvideo.configs.sample.wan.SelfForcingWan2_2_T2V_A14B_480P_SamplingParam
|
||||
- fastvideo.configs.sample.wan.Wan2_2_I2V_A14B_SamplingParam
|
||||
- fastvideo.configs.sample.wan.Wan2_2_T2V_A14B_SamplingParam
|
||||
profile_owned:
|
||||
action_list:
|
||||
target: request.extensions.hunyuangamecraft.action_list
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
action_speed_list:
|
||||
target: request.extensions.hunyuangamecraft.action_speed_list
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
camera_states:
|
||||
target: request.extensions.hunyuangamecraft.camera_states
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
camera_trajectory:
|
||||
target: request.extensions.hunyuangamecraft.camera_trajectory
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
conditioning_mask:
|
||||
target: request.extensions.hunyuangamecraft.conditioning_mask
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
gt_latents:
|
||||
target: request.extensions.hunyuangamecraft.gt_latents
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
prompt_attention_mask:
|
||||
target: request.extensions.hyworld.prompt_attention_mask
|
||||
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
|
||||
negative_attention_mask:
|
||||
target: request.extensions.hyworld.negative_attention_mask
|
||||
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
|
||||
ltx2_cfg_scale_audio:
|
||||
target: request.extensions.ltx2.cfg_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_cfg_scale_video:
|
||||
target: request.extensions.ltx2.cfg_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_modality_scale_audio:
|
||||
target: request.extensions.ltx2.modality_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_modality_scale_video:
|
||||
target: request.extensions.ltx2.modality_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_rescale_scale:
|
||||
target: request.extensions.ltx2.rescale_scale
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_blocks_audio:
|
||||
target: request.extensions.ltx2.stg_blocks_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_blocks_video:
|
||||
target: request.extensions.ltx2.stg_blocks_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_scale_audio:
|
||||
target: request.extensions.ltx2.stg_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_scale_video:
|
||||
target: request.extensions.ltx2.stg_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
sampling_param_extensions: {}
|
||||
|
||||
openai_image_request:
|
||||
kept:
|
||||
|
||||
@@ -12,7 +12,7 @@ FastVideo maps a Diffusers-style repo into a pipeline like this:
|
||||
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
|
||||
weight name translation.
|
||||
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
|
||||
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
|
||||
- `fastvideo/api/sampling_param.py`: runtime sampling parameters.
|
||||
- `fastvideo/pipelines/basic/*`: end-to-end pipelines.
|
||||
- `fastvideo/pipelines/stages/*`: reusable pipeline stages.
|
||||
- `fastvideo/models/loader/*`: component loaders for Diffusers-style repos.
|
||||
@@ -26,7 +26,7 @@ Minimal usage (from `examples/inference/basic/basic.py`):
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
|
||||
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
|
||||
@@ -49,8 +49,9 @@ runtime parameters consistent:
|
||||
- `fastvideo/configs/models/`: architecture definitions, layer shapes, and
|
||||
`param_names_mapping` rules for key renaming.
|
||||
- `fastvideo/configs/pipelines/`: pipeline wiring and required components.
|
||||
- `fastvideo/configs/sample/`: default sampling parameters (steps, frames,
|
||||
guidance scale, resolution, fps).
|
||||
- `fastvideo/api/sampling_param.py`: sampling parameters (steps, frames,
|
||||
guidance scale, resolution, fps). Defaults come from profiles in
|
||||
`fastvideo/pipelines/basic/<family>/profiles.py`.
|
||||
- `fastvideo/registry.py`: unified registry for pipeline config + sampling
|
||||
defaults and model metadata resolution, defined via explicit
|
||||
`register_configs(...)` blocks (no separate dict registries).
|
||||
@@ -142,7 +143,7 @@ How this maps to FastVideo:
|
||||
- `T5TokenizerFast` -> loaded via HF in `fastvideo/models/loader/`
|
||||
- `UniPCMultistepScheduler` -> loaded via Diffusers scheduler utilities
|
||||
- Pipeline defaults -> `fastvideo/configs/pipelines/wan.py`
|
||||
- Sampling defaults -> `fastvideo/configs/sample/wan.py`
|
||||
- Sampling defaults -> `fastvideo/pipelines/basic/wan/profiles.py`
|
||||
|
||||
## Pipeline system
|
||||
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
# Streaming WebSocket Server Contract
|
||||
|
||||
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
|
||||
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
|
||||
document is the authoritative spec for the message catalogue and the
|
||||
session state machine. Any change to either must update this document
|
||||
in the same PR that touches `protocol.py` or `session.py`.
|
||||
|
||||
## Endpoint
|
||||
|
||||
| Path | Protocol | Purpose |
|
||||
|---|---|---|
|
||||
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
|
||||
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
|
||||
|
||||
The server is launched by `fastvideo serve --config <serve.yaml>` when
|
||||
the config carries a `streaming:` block. Without that block the same CLI
|
||||
launches the OpenAI stateless HTTP server instead.
|
||||
|
||||
## Connection lifecycle
|
||||
|
||||
Every WebSocket connection holds exactly one `Session`. Sessions move
|
||||
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
|
||||
|
||||
```
|
||||
┌──────────────┐
|
||||
│ INITIALIZING │ ← WebSocket accepted, before init frame
|
||||
└──────┬───────┘
|
||||
│ session_init_v2 received
|
||||
┌──────────────┼──────────────┐
|
||||
▼ ▼ ▼
|
||||
QUEUED GPU_BINDING REJECTED
|
||||
│ │ ↑
|
||||
│ slot ready │ │ max-sessions hit
|
||||
▼ ▼ │ or invalid init
|
||||
┌────────┐ │
|
||||
│ ACTIVE │ ────────┘
|
||||
└────┬───┘
|
||||
segment loop │
|
||||
│
|
||||
┌───────────┼───────────┐
|
||||
▼ ▼ ▼
|
||||
COMPLETE ERROR TIMEOUT
|
||||
(clean leave) (any failure) (idle / segment_cap reached)
|
||||
```
|
||||
|
||||
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
|
||||
no transitions out. The transition matrix is enforced in
|
||||
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
|
||||
|
||||
`SessionManager` enforces the per-process budgets pulled from
|
||||
`StreamingConfig`:
|
||||
|
||||
- `session_timeout_seconds` — idle reaper drops sessions that haven't
|
||||
advanced; non-terminal sessions transition to `TIMEOUT`.
|
||||
- `generation_segment_cap` — a session that hits the cap transitions to
|
||||
`COMPLETE` after the last segment ships.
|
||||
|
||||
## Message catalogue
|
||||
|
||||
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
|
||||
`protocol.py` are the source of truth; this table is the human-readable
|
||||
view.
|
||||
|
||||
### Client → server
|
||||
|
||||
| `type` | Required fields | Purpose |
|
||||
|---|---|---|
|
||||
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
|
||||
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
|
||||
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
|
||||
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
|
||||
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
|
||||
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
|
||||
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
|
||||
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
|
||||
|
||||
The opening frame must be `session_init_v2`. Any other first frame is
|
||||
rejected with an `error` (code `invalid_message`) and the WebSocket is
|
||||
closed.
|
||||
|
||||
### Server → client
|
||||
|
||||
| `type` | Carries | When emitted |
|
||||
|---|---|---|
|
||||
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
|
||||
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
|
||||
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
|
||||
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
|
||||
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
|
||||
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
|
||||
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
|
||||
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
|
||||
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
|
||||
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
|
||||
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
|
||||
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
|
||||
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
|
||||
|
||||
## Continuation state
|
||||
|
||||
The session optionally accepts a `continuation_state` dict inside the
|
||||
opening `session_init_v2` frame. When present, the server hydrates it
|
||||
into a `ContinuationState(kind, payload)` envelope and feeds it as the
|
||||
`request.state` on the first segment's `GenerationRequest` — letting a
|
||||
client resume after a disconnect, migrate sessions across processes,
|
||||
or replay a prior session.
|
||||
|
||||
After every segment, if the runtime returns a fresh state, the server
|
||||
persists it to the `SessionStore` so a `snapshot_state` request can
|
||||
export it. The store and serialization contracts live with the model
|
||||
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
|
||||
|
||||
## Example flow
|
||||
|
||||
```
|
||||
client server
|
||||
────── ──────
|
||||
WS /v1/stream ─────── connect ─────────────────────────►
|
||||
◄────── (accept)
|
||||
|
||||
{"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox in snow", "the fox jumps"],
|
||||
"initial_image": {...},
|
||||
"stream_mode": "av_fmp4"} ─────────────────────────────►
|
||||
|
||||
(validate, queue, bind)
|
||||
◄──── {"type": "queue_status",
|
||||
"position": 0, "queue_depth": 0}
|
||||
◄──── {"type": "gpu_assigned",
|
||||
"gpu_id": 0, "model_id": "..."}
|
||||
◄──── {"type": "ltx2_stream_start", ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "a fox in snow",
|
||||
"source": "curated"} ───────────────────────────────────►
|
||||
(run pipeline)
|
||||
◄──── {"type": "ltx2_segment_start",
|
||||
"segment_idx": 1, ...}
|
||||
◄──── {"type": "step_complete",
|
||||
"segment_idx": 1, "timings": {...}}
|
||||
◄──── {"type": "media_init",
|
||||
"segment_idx": 1,
|
||||
"mime": "video/mp4", ...}
|
||||
◄──── <binary fMP4 init segment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── {"type": "media_segment_complete",
|
||||
"segment_idx": 1, "chunks": 12}
|
||||
◄──── {"type": "ltx2_segment_complete",
|
||||
"segment_idx": 1, ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "the fox jumps"} ─────────────────────────────►
|
||||
(segment 2 …)
|
||||
|
||||
{"type": "snapshot_state"} ──────────────────────────────►
|
||||
◄──── {"type": "continuation_state_snapshot",
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, ...}}
|
||||
|
||||
(close) ──────────────────────────────────────────────────►
|
||||
(session → COMPLETE)
|
||||
```
|
||||
|
||||
## Backward / forward compatibility
|
||||
|
||||
- Adding a new client message: append a Pydantic model to `protocol.py`
|
||||
with a unique `type`; add the discriminator entry to `ClientMessage`;
|
||||
add a row to the table above. Old clients that don't send the new
|
||||
message remain compatible.
|
||||
- Adding a new server message: emit only when a new feature flag is
|
||||
enabled (or always emit, since clients ignore unknown types).
|
||||
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
|
||||
→ `session_init_v3`) and accept both for one release cycle. Never
|
||||
silently change field semantics under the same `type`.
|
||||
@@ -243,7 +243,7 @@ for step in range(start_step, max_steps):
|
||||
|
||||
```bash
|
||||
# Install
|
||||
uv pip install -e .[dev]
|
||||
uv pip install -e ".[dev]"
|
||||
|
||||
# Run DMD2 distillation on Wan 2.1
|
||||
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
|
||||
|
||||
@@ -86,3 +86,25 @@ sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
|
||||
- Learning rate: 2e-5
|
||||
- Training steps: 3000 (~12 hours)
|
||||
- HSDP shard dim: 1
|
||||
|
||||
## 🧭 Note on `real_score_guidance_scale`
|
||||
|
||||
The teacher CFG used inside the DMD loss follows the DMD2 reference
|
||||
implementation and uses the parameterization
|
||||
|
||||
```
|
||||
x = x_cond + w * (x_cond - x_uncond)
|
||||
```
|
||||
|
||||
rather than the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)`. The
|
||||
two are mathematically equivalent up to a constant offset:
|
||||
|
||||
| `real_score_guidance_scale` (`w`) | Equivalent standard CFG (`w + 1`) | Output |
|
||||
|-----------------------------------|-----------------------------------|-----------------------|
|
||||
| `-1` | `0` | unconditional |
|
||||
| `0` | `1` | conditional |
|
||||
| `3.5` (default) | `4.5` | strong guidance |
|
||||
|
||||
So `real_score_guidance_scale` should be read as the **extra** guidance
|
||||
strength added on top of the conditional prediction. When porting values
|
||||
from a paper that uses the Ho & Salimans form, subtract 1.
|
||||
|
||||
@@ -27,7 +27,7 @@ uv pip install fastvideo
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### From source
|
||||
@@ -41,11 +41,11 @@ uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
Alternative with Conda environment:
|
||||
Alternative with Conda environment (still drives installs through `uv`):
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
@@ -58,14 +58,16 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install FlashAttention:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -87,7 +89,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
@@ -101,7 +103,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
|
||||
@@ -57,8 +57,10 @@ uv pip install fastvideo
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -80,7 +82,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -33,7 +33,7 @@ The following two classes `PipelineConfig` and `SamplingParam` are used to confi
|
||||
|
||||
### SamplingParam
|
||||
|
||||
::: fastvideo.configs.sample.base.SamplingParam
|
||||
::: fastvideo.api.sampling_param.SamplingParam
|
||||
options:
|
||||
show_root_heading: true
|
||||
show_source: false
|
||||
|
||||
@@ -128,19 +128,14 @@ Concrete hierarchy: `DiTConfig` → `DiTArchConfig`, `VAEConfig` →
|
||||
- `dump_to_json()` / `load_from_json()` — JSON persistence. Callable
|
||||
fields and `arch_config` are excluded from dumps.
|
||||
|
||||
### SamplingParam (`fastvideo/configs/sample/`)
|
||||
### SamplingParam (`fastvideo/api/sampling_param.py`)
|
||||
|
||||
Generation parameters separate from pipeline config. Each model family
|
||||
provides defaults:
|
||||
provides defaults via a profile (see `fastvideo/pipelines/basic/<family>/profiles.py`):
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class WanT2V_1_3B_SamplingParam(SamplingParam):
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
guidance_scale: float = 3.0
|
||||
num_inference_steps: int = 50
|
||||
sp = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sp.height == 480, sp.width == 832, sp.num_frames == 81, etc.
|
||||
```
|
||||
|
||||
## Component Loading
|
||||
@@ -430,9 +425,9 @@ User: generator.generate_video(prompt, ...)
|
||||
`fastvideo/configs/pipelines/<model>.py`. Set DiT/VAE/encoder configs,
|
||||
flow_shift, precision defaults.
|
||||
|
||||
2. **Sampling param** — Create a `SamplingParam` subclass in
|
||||
`fastvideo/configs/sample/<model>.py`. Set default height, width,
|
||||
num_frames, guidance_scale, num_inference_steps.
|
||||
2. **Sampling param profile** — Create a profile in
|
||||
`fastvideo/pipelines/basic/<model>/profiles.py` with default height,
|
||||
width, num_frames, guidance_scale, num_inference_steps.
|
||||
|
||||
3. **Register configs** — In `fastvideo/registry.py`, add a
|
||||
`register_configs()` call inside `_register_configs()` with
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
- Install MoGe:
|
||||
|
||||
```bash
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
```
|
||||
|
||||
- If you hit `ImportError: libGL.so.1` (common on Ubuntu/headless nodes), you can try installing OpenCV runtime libs:
|
||||
@@ -89,7 +89,7 @@ GEN3C defaults in FastVideo:
|
||||
|
||||
These values are defined in:
|
||||
|
||||
- `fastvideo/configs/sample/gen3c.py`
|
||||
- `fastvideo/pipelines/basic/gen3c/profiles.py`
|
||||
- `fastvideo/configs/pipelines/gen3c.py`
|
||||
|
||||
and align with the official GEN3C inference defaults in:
|
||||
|
||||
@@ -54,7 +54,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
We recommend always installing [Flash Attention 2](https://github.com/Dao-AILab/flash-attention):
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
uv pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://github.com/Dao-AILab/flash-attention?tab=readme-ov-file#flashattention-3-beta-release) by compiling it from source (takes about 10 minutes for me):
|
||||
@@ -63,7 +63,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
|
||||
|
||||
cd hopper
|
||||
pip install ninja
|
||||
uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
@@ -98,7 +98,7 @@ To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please co
|
||||
```bash
|
||||
git clone https://github.com/thu-ml/SageAttention.git
|
||||
cd sageattention
|
||||
python setup.py install # or pip install -e .
|
||||
python setup.py install # or uv pip install -e .
|
||||
```
|
||||
|
||||
### Sage Attention 3
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
uv pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
uv pip install vsa
|
||||
```
|
||||
|
||||
### Data-free Distillation
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
uv pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
def main():
|
||||
|
||||
@@ -7,7 +7,7 @@ and the GEN3C diffusion model.
|
||||
|
||||
Requirements:
|
||||
1. Install MoGe:
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
If you hit `ImportError: libGL.so.1`, install:
|
||||
sudo apt-get update && sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1
|
||||
2. Download and convert weights:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15_1080p"
|
||||
def main():
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
OUTPUT_PATH = "video_samples_lingbotworld"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal user-runnable example for the daVinci-MagiHuman base AV pipeline.
|
||||
|
||||
Produces an mp4 with both video (Wan 2.2 TI2V-5B VAE) and audio (Stable
|
||||
Audio Open 1.0 VAE, first-class FastVideo port in
|
||||
`fastvideo/models/vaes/oobleck.py`) muxed together via PyAV.
|
||||
|
||||
Prerequisites (one-off):
|
||||
|
||||
# Accept terms of use on the gated HF repos with your HF_TOKEN:
|
||||
# - https://huggingface.co/google/t5gemma-9b-9b-ul2
|
||||
# - https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
# All four cross-variant shared components (Wan 2.2 VAE, T5-Gemma
|
||||
# encoder + tokenizer, Stable Audio VAE) are lazy-loaded from their
|
||||
# canonical upstream HF repos on first build, so a single ~25 GB
|
||||
# cache is shared across every MagiHuman variant.
|
||||
|
||||
The umbrella HF repo `FastVideo/MagiHuman-Diffusers` holds all four
|
||||
variants (base / distill / sr_540p / sr_1080p) under sibling subfolders
|
||||
and FastVideo will download just the requested subfolder. Local
|
||||
conversion via `scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py`
|
||||
is also supported.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/base",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_video/magi_human_basic/output_magi_human.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Defaults pulled from the registered preset (magi_human_base):
|
||||
# height=256, width=448, fps=25, num_inference_steps=32, seed=42.
|
||||
# Override here only if you have a specific QA scenario.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal user-runnable example for the daVinci-MagiHuman DMD-2 distilled
|
||||
text-to-AV pipeline.
|
||||
|
||||
Same arch as the base model (`basic_magi_human.py`) but with DMD-2 distilled
|
||||
weights: 8 denoising steps, no classifier-free guidance. ~4x faster than
|
||||
base at the same 256x480 resolution. Mirrors upstream
|
||||
`daVinci-MagiHuman/example/distill/run_T2V.sh`.
|
||||
|
||||
Prerequisites (one-off):
|
||||
|
||||
# 1) Accept terms on the gated HF repos with your HF_TOKEN:
|
||||
# - https://huggingface.co/google/t5gemma-9b-9b-ul2
|
||||
# - https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
# Cross-variant shared components (Wan 2.2 VAE + T5-Gemma + Stable
|
||||
# Audio VAE) are lazy-loaded from their canonical upstream HF repos
|
||||
# and shared with the base variant cache.
|
||||
# 2) Convert the distill subfolder of GAIR/daVinci-MagiHuman:
|
||||
python scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py \\
|
||||
--source GAIR/daVinci-MagiHuman \\
|
||||
--subfolder distill \\
|
||||
--output converted_weights/magi_human_distill \\
|
||||
--cast-bf16
|
||||
# `--cast-bf16` is recommended (61 GB fp32 -> 30 GB bf16); the FV pipeline
|
||||
# loads bf16 anyway, and the conversion keeps norms / RoPE bands fp32.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/distill",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_video/magi_human_basic/output_magi_human_distill.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Defaults pulled from the registered preset (magi_human_distill):
|
||||
# height=256, width=480, fps=25, num_inference_steps=8, cfg=1, seed=42.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal daVinci-MagiHuman DMD-2 distilled text+image-to-AV example."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanDistillI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/distill",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanI2VPipeline",
|
||||
pipeline_config=MagiHumanDistillI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_distill_ti2v/output_magi_human_distill_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-1080p text-to-AV in FastVideo.
|
||||
|
||||
Build the converted repo on large local storage, then symlink it into the
|
||||
workspace:
|
||||
|
||||
python scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py \
|
||||
--source GAIR/daVinci-MagiHuman \
|
||||
--subfolder base \
|
||||
--sr-source GAIR/daVinci-MagiHuman \
|
||||
--sr-subfolder 1080p_sr \
|
||||
--output /raid/william5lin_converted_weights/magi_human_sr_1080p \
|
||||
--cast-bf16
|
||||
ln -s /raid/william5lin_converted_weights/magi_human_sr_1080p \
|
||||
converted_weights/magi_human_sr_1080p
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanSR1080pConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_1080p",
|
||||
num_gpus=1,
|
||||
override_pipeline_cls_name="MagiHumanSR1080pPipeline",
|
||||
pipeline_config=MagiHumanSR1080pConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_video/magi_human_sr1080p/output_magi_human_sr1080p.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-1080p text+image-to-AV in FastVideo."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanSR1080pI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_1080p",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanSR1080pI2VPipeline",
|
||||
pipeline_config=MagiHumanSR1080pI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_sr1080p_ti2v/output_magi_human_sr1080p_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-540p text-to-AV in FastVideo.
|
||||
|
||||
The converted repo must contain both ``transformer/`` (base DiT) and
|
||||
``sr_transformer/`` (540p SR DiT). Build it with:
|
||||
|
||||
python scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py \
|
||||
--source GAIR/daVinci-MagiHuman \
|
||||
--subfolder base \
|
||||
--sr-source GAIR/daVinci-MagiHuman \
|
||||
--sr-subfolder 540p_sr \
|
||||
--output converted_weights/magi_human_sr_540p
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_540p",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_video/magi_human_sr540p/output_magi_human_sr540p.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-540p text+image-to-AV in FastVideo."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanSR540pI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_540p",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanSRI2VPipeline",
|
||||
pipeline_config=MagiHumanSR540pI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_sr540p_ti2v/output_magi_human_sr540p_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal daVinci-MagiHuman base text+image-to-AV example."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanBaseI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/base",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanI2VPipeline",
|
||||
pipeline_config=MagiHumanBaseI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_ti2v/output_magi_human_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
def main():
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — text-to-audio (baseline) example.
|
||||
|
||||
User story (game-audio designer, prototyping):
|
||||
"I'm prototyping a level and I need 6 seconds of background
|
||||
ambience — gentle wind, distant thunder, a hint of birdsong. I
|
||||
don't want to dig through a sound library; I want to type what I
|
||||
hear in my head and get a wav back. If it's wrong I'll iterate
|
||||
on the prompt. This is the first stop."
|
||||
|
||||
User story (musician sketching ideas):
|
||||
"I want to bounce a 30s lo-fi drum loop to use as a placeholder
|
||||
bed while I build the rest of the track. Type prompt, get audio,
|
||||
drop into the DAW. The actual production beat I'll record
|
||||
myself, but I need *something* to write the chords against."
|
||||
|
||||
User story (researcher exploring the model):
|
||||
"First time touching Stable Audio Open — what does it sound
|
||||
like at default settings? This is the smallest amount of code
|
||||
that goes from prompt to mp4."
|
||||
|
||||
How it works:
|
||||
Pure text-to-audio (T2A). The pipeline runs:
|
||||
T5 + NumberConditioner -> StableAudioDiT -> Oobleck VAE
|
||||
via the `dpmpp-3m-sde` k-diffusion sampler. All components are
|
||||
FastVideo-native — no diffusers / transformers model imports at
|
||||
runtime (see REVIEW item 30). Mirrors upstream
|
||||
`stable_audio_tools.inference.generation.generate_diffusion_cond`
|
||||
bit-for-bit (~0.2% abs_mean drift on 25 steps).
|
||||
|
||||
Tunable knobs (the "creative dials"):
|
||||
audio_end_in_s
|
||||
1–6 — quick ideation (sub-10s wall clock at 100 steps)
|
||||
10–30 — full musical phrase / loop length (the README example
|
||||
uses 30s)
|
||||
47.5 — model maximum (full sample_size = 2097152 / 44100 Hz)
|
||||
num_inference_steps
|
||||
25 — fast preview, occasional artifacts
|
||||
100 — preset default (matches the HF model card)
|
||||
250 — diminishing returns past here
|
||||
guidance_scale
|
||||
3 — looser, more variation per seed
|
||||
7 — preset default; matches README
|
||||
12+ — sharper but can sound "fried"
|
||||
|
||||
Prerequisites:
|
||||
1. Accept the terms on https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
and export your HF token in the shell:
|
||||
export HF_TOKEN=hf_...
|
||||
2. Install optional inference deps (one-time):
|
||||
uv pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_audio/stable_audio_basic/output_stable_audio.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# 6-second clip; the model max is ~47.5s.
|
||||
audio_end_in_s=6.0,
|
||||
# The registered preset gives 100 steps + CFG=7.0 by default;
|
||||
# override num_inference_steps / guidance_scale here for QA.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,77 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — audio-to-audio variation example.
|
||||
|
||||
User story (musician, late at night):
|
||||
"I generated this 12-second lo-fi loop earlier and I love the chord
|
||||
progression and overall vibe, but the snare hit at 0:08 sounds wrong
|
||||
and the rhythm feels stiff. I don't want to start over from scratch
|
||||
and lose what's working — I want the model to keep the harmony and
|
||||
mood but reroll the percussion + groove."
|
||||
|
||||
User story (sound designer, on a deadline):
|
||||
"I have one good 'sword clang' SFX. The art director wants 8 sibling
|
||||
variations that all feel like the same sword from different angles —
|
||||
same metal, same weight, slightly different impact. I'd rather
|
||||
refine my one good take than text-prompt my way through 50 misses."
|
||||
|
||||
Pass `init_audio=path/to/clip` (any wav/mp3/mp4/m4a/flac the standard
|
||||
deps decode) and the model will use it as a starting point for the
|
||||
text prompt instead of pure noise.
|
||||
|
||||
Picking `init_audio_strength` (0.0 to 1.0):
|
||||
|
||||
Higher = closer to the source clip. Lower = more transformation.
|
||||
(Same convention as the "Input Audio Strength" slider in
|
||||
Stability's commercial Stable Audio web UI, so values transfer
|
||||
directly.)
|
||||
|
||||
| strength | what you get |
|
||||
|----------|----------------------------------------------------|
|
||||
| 1.00 | Output ≈ reference. No transformation. |
|
||||
| 0.85 | Texture micro-variation only. |
|
||||
| 0.70 | Light reroll, same instruments. |
|
||||
| 0.60 | Default. Instrument identity is replaceable |
|
||||
| | (cello can take over from piano on the same notes).|
|
||||
| 0.50 | Heavy — only melody / chord progression survives. |
|
||||
| 0.30 | Reference acts as a loose mood prompt. |
|
||||
| 0.00 | Plain T2A — reference ignored. |
|
||||
|
||||
Rule of thumb by intent:
|
||||
* "Fix one part of this clip" -> 0.75 .. 0.85
|
||||
* "Same notes, different instrument" -> 0.55 .. 0.65
|
||||
* "Same chord progression, new content" -> 0.40 .. 0.55
|
||||
* "Use this as a loose mood prompt" -> 0.20 .. 0.35
|
||||
|
||||
If the reference timbre is bleeding through more than you want,
|
||||
lower it; if the structure is gone, raise it.
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Change the piano to a cello playing the same notes"
|
||||
# Path to any audio-bearing file (wav, mp3, mp4, m4a, flac, ...).
|
||||
# Set to `None` to skip A2A and run plain T2A.
|
||||
INIT_AUDIO_PATH: str | None = None
|
||||
# Reference fidelity in [0, 1] -- higher = closer to source.
|
||||
INIT_AUDIO_STRENGTH = 0.6
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_a2a/output_a2a.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=6.0,
|
||||
init_audio=INIT_AUDIO_PATH,
|
||||
init_audio_strength=INIT_AUDIO_STRENGTH,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,84 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — inpainting / outpainting (loop extension) example.
|
||||
|
||||
User story (loop extension — the killer app):
|
||||
"I have a 6-second drum loop my client likes. They want it as
|
||||
background bed for a 30-second ad. I need it to loop seamlessly,
|
||||
but a hard cut every 6s sounds bad. Let me extend it to 30s,
|
||||
keeping the first 6s exactly as-is and letting the model continue
|
||||
the groove for the remaining 24s."
|
||||
|
||||
User story (audio repair):
|
||||
"There's a microphone bump at 0:14 in this 30-second field
|
||||
recording — really obvious in headphones. Mask out 0:13 to 0:15
|
||||
and let the model regenerate plausible ambience that blends in.
|
||||
Everything else stays exactly as I recorded it."
|
||||
|
||||
User story (transition smoothing):
|
||||
"I have two 10-second clips I want to crossfade. Mask out a 1s
|
||||
overlap region in the middle and let the model invent a coherent
|
||||
transition between the two."
|
||||
|
||||
How it works (RePaint-style blending):
|
||||
Stable Audio Open 1.0 wasn't trained as an inpainting model
|
||||
(`model_type=diffusion_cond`, not `diffusion_cond_inpaint`), so we
|
||||
can't use the upstream's mask-conditioned approach directly. We
|
||||
use the RePaint trick instead, which works on any v-prediction
|
||||
diffusion model:
|
||||
|
||||
1. Encode the reference clip into latent space.
|
||||
2. At every denoising step `i`, replace the kept region of the
|
||||
in-flight latent (where mask == 1) with the reference
|
||||
re-noised to the next timestep's sigma. Only the unkept
|
||||
region (mask == 0) is freely denoised.
|
||||
3. After the loop, the kept region is exactly the reference;
|
||||
the unkept region is freshly generated content.
|
||||
|
||||
This is approximate compared to a properly trained inpainting
|
||||
checkpoint — the seam between kept/unkept can have slight EQ
|
||||
discontinuity — but it works on the existing public model.
|
||||
|
||||
Tunable: the mask is a 1-D tensor in {0, 1} at the model's sample
|
||||
rate. Conventions:
|
||||
1.0 = keep this sample from the reference
|
||||
0.0 = regenerate this sample
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
|
||||
# Required: path to the reference audio file (wav, mp3, mp4, m4a, flac,
|
||||
# ...) you want to extend or repair. The pipeline raises if a mask is
|
||||
# passed without a reference, so this must be a real path.
|
||||
REFERENCE_AUDIO_PATH = "path/to/your/loop.wav"
|
||||
KEEP_SECONDS = 6.0 # first KEEP_SECONDS preserved exactly
|
||||
TOTAL_SECONDS = 12.0 # extend the loop to this duration
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not os.path.isfile(REFERENCE_AUDIO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
|
||||
"Edit this script to point at a real audio file (wav/mp3/mp4/"
|
||||
"m4a/flac) before running.")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_inpaint/output_inpaint.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=TOTAL_SECONDS,
|
||||
inpaint_audio=REFERENCE_AUDIO_PATH,
|
||||
# Tuple form: keep first KEEP_SECONDS, regenerate the rest.
|
||||
inpaint_mask=(KEEP_SECONDS, TOTAL_SECONDS),
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open Small — fast / lightweight T2A example.
|
||||
|
||||
User story (interactive UI builder):
|
||||
"I'm building a sound-design UI where the user types a prompt and
|
||||
we want sub-2-second feedback so the experience feels like
|
||||
autocomplete, not a render queue. The full Stable Audio Open 1.0
|
||||
takes ~8s on a single GPU; the small variant takes a fraction of
|
||||
that — quality is lower but completely usable for real-time
|
||||
iteration."
|
||||
|
||||
User story (overnight batch jobs):
|
||||
"I'm generating 10,000 short SFX variants for a procedural game.
|
||||
Wall-clock matters more than per-clip polish — give me the small
|
||||
model so I can fit the run in one night instead of a week."
|
||||
|
||||
How it works:
|
||||
The small variant is a separate Stability AI checkpoint
|
||||
(`stabilityai/stable-audio-open-small`) that ships the same Oobleck
|
||||
VAE as the 1.0 base model but a smaller / faster DiT (`embed_dim=1024`,
|
||||
`depth=16`, `qk_norm="ln"`) and only one duration conditioner
|
||||
(`seconds_total`, no `seconds_start`). FastVideo loads from the
|
||||
converted Diffusers-format repo `FastVideo/stable-audio-open-small-Diffusers`
|
||||
via the standard component loader; per-variant arch fields come
|
||||
from `transformer/config.json` and `conditioner/config.json`.
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`. The converted repo is
|
||||
public so no gated-access flow is required.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-small-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_audio/stable_audio_small/output_stable_audio_small.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Small variant trains on a ~11.9s window — keep `audio_end_in_s`
|
||||
# at or below that.
|
||||
audio_end_in_s=6.0,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_1_Fun"
|
||||
OUTPUT_NAME = "wan2.1_test"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
|
||||
def main():
|
||||
|
||||
@@ -5,7 +5,7 @@ import time
|
||||
|
||||
import gradio as gr
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from copy import deepcopy
|
||||
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ import tempfile
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
MODEL_PATH_MAPPING = {
|
||||
|
||||
@@ -185,7 +185,7 @@ class BaseModelDeployment:
|
||||
|
||||
def _initialize_generator(self, config: Dict[str, Any]) -> None:
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
print(f"Initializing model: {self.model_path}")
|
||||
self.generator = VideoGenerator.from_pretrained(
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
Inference using a LoRA checkpoint from FastVideo trainer.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# Cosmos Predict2 2B T2V finetune config.
|
||||
#
|
||||
# Data must be preprocessed with Cosmos VAE + T5 text encoder
|
||||
# into parquet format before training.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: nvidia/Cosmos-Predict2-2B-Video2World
|
||||
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: 8
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos_preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
# Cosmos VAE: 4x temporal, 8x spatial compression.
|
||||
# 93 frames -> 24 latent frames, 480x832 -> 60x104
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 5000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos_finetune
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos
|
||||
run_name: cosmos_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
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.cosmos.cosmos_pipeline.Cosmos2VideoToWorldPipeline
|
||||
dataset_file: data/cosmos_preprocessed/validation_prompts.json
|
||||
every_steps: 100
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -0,0 +1,79 @@
|
||||
# Cosmos-Predict2.5-2B Text-to-World overfitting test config.
|
||||
#
|
||||
# Overfits on a few short videos (480x832, 93 frames) to verify the
|
||||
# Cosmos 2.5 training plugin works end-to-end.
|
||||
#
|
||||
# Preprocess data first:
|
||||
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# bash examples/train/run.sh examples/train/configs/overfit_cosmos25_t2w.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: KyleShao/Cosmos-Predict2.5-2B-Diffusers
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
flow_shift: 1.0
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos25_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos25_overfit
|
||||
training_state_checkpointing_steps: 50
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos25
|
||||
run_name: cosmos25_overfit
|
||||
|
||||
model:
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
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.cosmos.cosmos2_5_pipeline.Cosmos2_5Pipeline
|
||||
dataset_file: data/cosmos25_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 150
|
||||
sampling_steps: [35]
|
||||
guidance_scale: 7.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -1,7 +1,7 @@
|
||||
cmake_minimum_required(VERSION 3.26 FATAL_ERROR)
|
||||
project(fastvideo-kernel LANGUAGES CXX)
|
||||
|
||||
# Prefer environment variable (used by CI or pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
# Prefer environment variable (used by CI or uv pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
|
||||
@@ -61,7 +61,16 @@ has_cmake_arg() {
|
||||
}
|
||||
|
||||
detect_with_torch() {
|
||||
uv run --active --no-project python -c "import torch
|
||||
# Prefer the active venv's python directly over `uv run --active --no-project`,
|
||||
# which on some uv versions provisions its own interpreter and misses packages
|
||||
# installed into VIRTUAL_ENV.
|
||||
local py
|
||||
if [[ -n "${VIRTUAL_ENV:-}" && -x "${VIRTUAL_ENV}/bin/python" ]]; then
|
||||
py="${VIRTUAL_ENV}/bin/python"
|
||||
else
|
||||
py="$(command -v python3 || command -v python)"
|
||||
fi
|
||||
"${py}" -c "import torch
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError('torch.cuda.is_available() is false')
|
||||
mj, mn = torch.cuda.get_device_capability(0)
|
||||
|
||||
@@ -5,6 +5,11 @@ from fastvideo_kernel.ops import (
|
||||
video_sparse_attn,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn import (
|
||||
block_sparse_attn,
|
||||
block_sparse_attn_from_indices,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.vmoba import (
|
||||
moba_attn_varlen,
|
||||
process_moba_input,
|
||||
@@ -22,6 +27,8 @@ from fastvideo_kernel.turbodiffusion_ops import (
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"block_sparse_attn",
|
||||
"block_sparse_attn_from_indices",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
@@ -6,6 +8,11 @@ from typing import Tuple
|
||||
import torch
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Backend selection helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_sm90_ops():
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
|
||||
@@ -25,38 +32,66 @@ def _is_sm90() -> bool:
|
||||
|
||||
|
||||
def _force_triton() -> bool:
|
||||
# Force Triton even on SM90 and even if the compiled extension is available.
|
||||
# Useful for CI / debugging / parity testing.
|
||||
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
|
||||
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Preferred map->index conversion used by the wrapper.
|
||||
# ---------------------------------------------------------------------------
|
||||
# Index helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
This wrapper **requires** the Triton implementation.
|
||||
If Triton (or the Triton map_to_index module) is not available, it raises.
|
||||
"""
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Compact a bool block_map to (q2k_idx, q2k_num). Legacy path only."""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
if block_map.dim() != 4:
|
||||
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
|
||||
raise ValueError(
|
||||
f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), "
|
||||
f"got shape={tuple(block_map.shape)}"
|
||||
)
|
||||
if block_map.dtype != torch.bool:
|
||||
block_map = block_map.to(torch.bool)
|
||||
|
||||
if not block_map.is_cuda:
|
||||
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
|
||||
raise RuntimeError(
|
||||
"block_map must be a CUDA tensor (Triton map_to_index required)."
|
||||
)
|
||||
|
||||
try:
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
|
||||
except Exception as e:
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index
|
||||
except Exception as e: # pragma: no cover - environment issue
|
||||
raise ImportError(
|
||||
"Triton map_to_index is required but not available. "
|
||||
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
|
||||
"Ensure Triton is installed and "
|
||||
"fastvideo_kernel.triton_kernels.index is importable."
|
||||
) from e
|
||||
return triton_map_to_index(block_map)
|
||||
|
||||
|
||||
def _invert_indices_for_backward(
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from fastvideo_kernel.triton_kernels.index import invert_indices
|
||||
return invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
|
||||
|
||||
|
||||
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
|
||||
"""Return `t` as a contiguous int32 tensor, raising a clear error on CPU input."""
|
||||
if not t.is_cuda:
|
||||
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
|
||||
if t.dtype != torch.int32:
|
||||
t = t.to(torch.int32)
|
||||
if not t.is_contiguous():
|
||||
t = t.contiguous()
|
||||
return t
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Triton backend custom ops (index-native)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_triton",
|
||||
mutates_args=(),
|
||||
@@ -66,34 +101,40 @@ def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
|
||||
triton_block_sparse_attn_forward,
|
||||
)
|
||||
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
o, M = triton_block_sparse_attn_forward(
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
return o, M
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
M = torch.empty(
|
||||
(q.shape[0], q.shape[1], q.shape[2]),
|
||||
device=q.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
return o, M
|
||||
|
||||
|
||||
@@ -109,20 +150,32 @@ def block_sparse_attn_backward_triton(
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output = grad_output.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
|
||||
triton_block_sparse_attn_backward,
|
||||
)
|
||||
|
||||
num_kv_blocks = int(variable_block_sizes.numel())
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(
|
||||
q2k_idx, q2k_num, num_kv_blocks
|
||||
)
|
||||
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
|
||||
# o/M are kernel outputs so are already contiguous.
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(
|
||||
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
|
||||
grad_output.contiguous(),
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
o,
|
||||
M,
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
return dq, dk, dv
|
||||
|
||||
@@ -135,7 +188,8 @@ def _block_sparse_attn_backward_triton_fake(
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q)
|
||||
@@ -144,19 +198,28 @@ def _block_sparse_attn_backward_triton_fake(
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _backward_triton(ctx, grad_o, grad_M):
|
||||
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _setup_context_triton(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
|
||||
o, M = output
|
||||
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
|
||||
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
|
||||
def _backward_triton(ctx, grad_o, grad_M):
|
||||
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(
|
||||
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None, None
|
||||
|
||||
|
||||
block_sparse_attn_triton.register_autograd(
|
||||
_backward_triton, setup_context=_setup_context_triton
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SM90 backend custom ops (index-native)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
@@ -168,21 +231,21 @@ def block_sparse_attn_sm90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
block_sparse_fwd, _ = _get_sm90_ops()
|
||||
if block_sparse_fwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
|
||||
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
o_padded, lse_padded = block_sparse_fwd(
|
||||
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
|
||||
q_padded.contiguous(),
|
||||
k_padded.contiguous(),
|
||||
v_padded.contiguous(),
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
return o_padded, lse_padded
|
||||
|
||||
@@ -192,11 +255,16 @@ def _block_sparse_attn_sm90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q_padded)
|
||||
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
|
||||
lse = torch.empty(
|
||||
(q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1),
|
||||
device=q_padded.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
return o, lse
|
||||
|
||||
|
||||
@@ -212,30 +280,34 @@ def block_sparse_attn_backward_sm90(
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
_, block_sparse_bwd = _get_sm90_ops()
|
||||
if block_sparse_bwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
|
||||
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
num_kv_blocks = int(variable_block_sizes.numel())
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(
|
||||
q2k_idx, q2k_num, num_kv_blocks
|
||||
)
|
||||
|
||||
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
|
||||
dq, dk, dv = block_sparse_bwd(
|
||||
q_padded,
|
||||
k_padded,
|
||||
v_padded,
|
||||
q_padded.contiguous(),
|
||||
k_padded.contiguous(),
|
||||
v_padded.contiguous(),
|
||||
o_padded,
|
||||
lse_padded,
|
||||
grad_output_padded,
|
||||
grad_output_padded.contiguous(),
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes.int(),
|
||||
variable_block_sizes,
|
||||
)
|
||||
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
|
||||
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
|
||||
# C++ kernel returns fp32 grads; cast back to the input dtype.
|
||||
out_dtype = grad_output_padded.dtype
|
||||
return dq.to(out_dtype), dk.to(out_dtype), dv.to(out_dtype)
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
|
||||
@@ -246,7 +318,8 @@ def _block_sparse_attn_backward_sm90_fake(
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q_padded)
|
||||
@@ -255,21 +328,57 @@ def _block_sparse_attn_backward_sm90_fake(
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _backward_sm90(ctx, grad_o, grad_lse):
|
||||
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_sm90(
|
||||
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _setup_context_sm90(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
|
||||
o, lse = output
|
||||
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
|
||||
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
|
||||
def _backward_sm90(ctx, grad_o, grad_lse):
|
||||
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_sm90(
|
||||
grad_o, q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None, None
|
||||
|
||||
|
||||
block_sparse_attn_sm90.register_autograd(
|
||||
_backward_sm90, setup_context=_setup_context_sm90
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def block_sparse_attn_from_indices(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Block-sparse attention with autograd, taking compact per-row KV indices."""
|
||||
# Normalize index tensors once at the public boundary so the custom ops
|
||||
# and their fakes can assume int32/contiguous. No-op on well-formed input.
|
||||
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
|
||||
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
|
||||
variable_block_sizes = _as_int32_contig(variable_block_sizes, "variable_block_sizes")
|
||||
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
use_sm90 = (
|
||||
(not _force_triton())
|
||||
and _is_sm90()
|
||||
and block_sparse_fwd is not None
|
||||
and block_sparse_bwd is not None
|
||||
)
|
||||
if use_sm90:
|
||||
return block_sparse_attn_sm90(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
|
||||
# to a multiple of the block size (64 tokens).
|
||||
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
def block_sparse_attn(
|
||||
@@ -279,16 +388,8 @@ def block_sparse_attn(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Unified block-sparse attention op with autograd support.
|
||||
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
|
||||
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
|
||||
"""
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
|
||||
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
|
||||
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
|
||||
# to a multiple of the block size (64 tokens).
|
||||
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
"""Bool-mask compat wrapper; prefer block_sparse_attn_from_indices."""
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
return block_sparse_attn_from_indices(
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import math
|
||||
import torch
|
||||
from .block_sparse_attn import block_sparse_attn
|
||||
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
|
||||
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
|
||||
|
||||
# Try to load the C++ extension
|
||||
@@ -125,13 +125,18 @@ def video_sparse_attn(
|
||||
out_c = out_c.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch, heads, q_seq_len, dim)
|
||||
|
||||
# Sparse branch
|
||||
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
|
||||
topk_idx = torch.topk(scores, topk, dim=-1).indices
|
||||
mask = torch.zeros_like(scores,
|
||||
dtype=torch.bool).scatter_(-1, topk_idx, True)
|
||||
|
||||
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
q2k_idx = topk_idx.to(torch.int32).contiguous()
|
||||
q2k_num = torch.full(
|
||||
(batch, heads, q_num_blocks),
|
||||
topk,
|
||||
dtype=torch.int32,
|
||||
device=q.device,
|
||||
)
|
||||
out_s = block_sparse_attn_from_indices(
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes
|
||||
)[0]
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
## pytorch sdpa version of block sparse ##
|
||||
from typing import Tuple
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
@@ -153,3 +154,114 @@ def map_to_index(block_map: torch.Tensor):
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _invert_indices_kernel(
|
||||
q2k_idx_ptr,
|
||||
q2k_num_ptr,
|
||||
k2q_idx_ptr,
|
||||
k2q_num_ptr,
|
||||
q2k_idx_b, q2k_idx_h, q2k_idx_q, q2k_idx_k,
|
||||
q2k_num_b, q2k_num_h, q2k_num_q,
|
||||
k2q_idx_b, k2q_idx_h, k2q_idx_k, k2q_idx_q,
|
||||
k2q_num_b, k2q_num_h, k2q_num_k,
|
||||
MAX_KV_PER_Q: tl.constexpr,
|
||||
):
|
||||
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
|
||||
pid_b = tl.program_id(0)
|
||||
pid_h = tl.program_id(1)
|
||||
pid_q = tl.program_id(2)
|
||||
|
||||
n = tl.load(
|
||||
q2k_num_ptr
|
||||
+ pid_b * q2k_num_b
|
||||
+ pid_h * q2k_num_h
|
||||
+ pid_q * q2k_num_q
|
||||
)
|
||||
|
||||
q2k_row = (
|
||||
q2k_idx_ptr
|
||||
+ pid_b * q2k_idx_b
|
||||
+ pid_h * q2k_idx_h
|
||||
+ pid_q * q2k_idx_q
|
||||
)
|
||||
|
||||
for i in tl.range(0, MAX_KV_PER_Q):
|
||||
if i < n:
|
||||
kv = tl.load(q2k_row + i * q2k_idx_k)
|
||||
count_ptr = (
|
||||
k2q_num_ptr
|
||||
+ pid_b * k2q_num_b
|
||||
+ pid_h * k2q_num_h
|
||||
+ kv * k2q_num_k
|
||||
)
|
||||
pos = tl.atomic_add(count_ptr, 1)
|
||||
tl.store(
|
||||
k2q_idx_ptr
|
||||
+ pid_b * k2q_idx_b
|
||||
+ pid_h * k2q_idx_h
|
||||
+ kv * k2q_idx_k
|
||||
+ pos * k2q_idx_q,
|
||||
pid_q,
|
||||
)
|
||||
|
||||
|
||||
def invert_indices(
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
|
||||
if q2k_idx.dim() != 4:
|
||||
raise ValueError(
|
||||
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
|
||||
)
|
||||
if q2k_num.dim() != 3:
|
||||
raise ValueError(
|
||||
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
|
||||
)
|
||||
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
|
||||
raise RuntimeError("invert_indices requires CUDA tensors.")
|
||||
|
||||
B, H, Nq, Mk = q2k_idx.shape
|
||||
if q2k_num.shape != (B, H, Nq):
|
||||
raise ValueError(
|
||||
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
|
||||
f"[B, H, Nq] = {(B, H, Nq)}"
|
||||
)
|
||||
|
||||
q2k_idx = q2k_idx.contiguous()
|
||||
q2k_num = q2k_num.contiguous()
|
||||
if q2k_idx.dtype != torch.int32:
|
||||
q2k_idx = q2k_idx.to(torch.int32)
|
||||
if q2k_num.dtype != torch.int32:
|
||||
q2k_num = q2k_num.to(torch.int32)
|
||||
|
||||
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
|
||||
# `Nq` is a tight upper bound on the compacted K->Q slots.
|
||||
k2q_idx = torch.empty(
|
||||
(B, H, num_kv_blocks, Nq),
|
||||
dtype=torch.int32,
|
||||
device=q2k_idx.device,
|
||||
)
|
||||
k2q_num = torch.zeros(
|
||||
(B, H, num_kv_blocks),
|
||||
dtype=torch.int32,
|
||||
device=q2k_idx.device,
|
||||
)
|
||||
|
||||
grid = (B, H, Nq)
|
||||
_invert_indices_kernel[grid](
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
q2k_idx.stride(0), q2k_idx.stride(1), q2k_idx.stride(2), q2k_idx.stride(3),
|
||||
q2k_num.stride(0), q2k_num.stride(1), q2k_num.stride(2),
|
||||
k2q_idx.stride(0), k2q_idx.stride(1), k2q_idx.stride(2), k2q_idx.stride(3),
|
||||
k2q_num.stride(0), k2q_num.stride(1), k2q_num.stride(2),
|
||||
MAX_KV_PER_Q=Mk,
|
||||
)
|
||||
|
||||
return k2q_idx, k2q_num
|
||||
|
||||
@@ -11,7 +11,7 @@ except ImportError:
|
||||
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError(
|
||||
"flash-attn is not installed. Please install it, e.g., `pip install flash-attn`."
|
||||
"flash-attn is not installed. Please install it, e.g., `uv pip install flash-attn`."
|
||||
)
|
||||
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# `fastvideo/` — Core Package
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Inference + training framework for video DiTs. Public API entry: `from fastvideo import VideoGenerator, PipelineConfig, SamplingParam`.
|
||||
|
||||
## Public Surface (`__init__.py`)
|
||||
|
||||
```python
|
||||
VideoGenerator # entrypoints/video_generator.py — high-level inference handle
|
||||
PipelineConfig # configs/pipelines/base.py — pipeline wiring dataclass
|
||||
SamplingParam # api/sampling_param.py — runtime sampling knobs
|
||||
```
|
||||
|
||||
CLI entry: `fastvideo` script → `entrypoints/cli/main.py` (subcommands: `generate`, `serve`, `bench`).
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
fastvideo/
|
||||
├── api/ # Schema + presets for the OpenAI-compatible serving layer
|
||||
├── attention/ # Backends + selector (FlashAttn / SageAttn / SDPA / VSA / VMoBA / SLA)
|
||||
├── configs/ # Per-model arch configs + per-pipeline configs (registry-driven)
|
||||
├── dataset/ # Dataloaders (pre-commit excluded — minimal lint surface)
|
||||
├── distributed/ # SP/TP groups, device communicators, init helpers
|
||||
├── entrypoints/ # cli/, openai/, streaming/, video_generator.py
|
||||
├── hooks/ # Runtime hook system for pipelines
|
||||
├── layers/ # Tensor-parallel linears + attention wrappers (port targets)
|
||||
├── models/ # DiT / VAE / encoder / scheduler / loader (pre-commit excluded)
|
||||
├── pipelines/ # basic/<model>/, preprocess/, stages/, training/
|
||||
├── platforms/ # CUDA/ROCm capability + AttentionBackendEnum
|
||||
├── third_party/ # Vendored externals (lint excluded; do not reformat)
|
||||
├── train/ # NEW modular trainer — methods × models × callbacks
|
||||
├── training/ # LEGACY monolithic *_training/distillation_pipeline.py
|
||||
├── worker/ # Multi-process / Ray executors
|
||||
├── workflow/ # Preprocessing workflow base class
|
||||
├── registry.py # Pipeline-config + model-class lookup (canonical)
|
||||
├── envs.py # Env-var declarations
|
||||
├── fastvideo_args.py# Runtime arg dataclass passed through pipelines
|
||||
└── utils.py # FlexibleArgumentParser, qualname resolver, etc.
|
||||
```
|
||||
|
||||
## Where to Look
|
||||
|
||||
| Task | Location |
|
||||
|------|----------|
|
||||
| Add a new pipeline class | `pipelines/basic/<model>/` + `configs/pipelines/<model>.py` + register in `registry.py` |
|
||||
| Add a new model component | `models/<role>/<model>.py` + `configs/models/<role>/<model>.py` |
|
||||
| Wire an existing model into a new pipeline | `pipelines/basic/<model>/presets.py` + reuse stages from `pipelines/stages/` |
|
||||
| Add a converter | `scripts/checkpoint_conversion/<model>_to_*.py` (separate dir, separate AGENTS.md) |
|
||||
| Add an attention backend | `attention/backends/<name>.py` + register in selector |
|
||||
| Add a runtime CLI flag | `fastvideo_args.py` (avoid `argparse` ad-hoc inside stages) |
|
||||
|
||||
## Conventions Specific Here
|
||||
|
||||
- `PipelineStage` subclasses (`pipelines/stages/`) own one verb each (encode, schedule, denoise, decode). Compose, don't fork.
|
||||
- Every pipeline reads from a `PipelineConfig` subclass and a `SamplingParam`. Never read raw env vars inside a stage — go through `fastvideo.envs`.
|
||||
- Logger setup: `from fastvideo.logger import init_logger; logger = init_logger(__name__)`. Do not call `logging.getLogger` directly.
|
||||
- Imports between `train/` and `training/` are **forbidden** — they are independent stacks.
|
||||
|
||||
## Pre-Commit Exclusions (do not assume linted)
|
||||
|
||||
These dirs are listed in `.pre-commit-config.yaml` `exclude`:
|
||||
|
||||
- `fastvideo/third_party/`, `fastvideo/dataset/`, `fastvideo/models/`
|
||||
|
||||
Editing files there will NOT trigger yapf/ruff/mypy/codespell. Format manually if a sibling file shows clear style; do not introduce new violations.
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
|
||||
@@ -7,21 +7,37 @@ from fastvideo.api.schema import (
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
GpuPoolConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
PlannedStage,
|
||||
PromptEnhancerConfig,
|
||||
PromptSafetyConfig,
|
||||
QuantizationConfig,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
ServerConfig,
|
||||
StreamingConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.presets import (
|
||||
InferencePreset,
|
||||
PresetStageSpec,
|
||||
get_all_preset_names,
|
||||
get_preset,
|
||||
get_presets_for_family,
|
||||
register_preset,
|
||||
validate_preset_selection,
|
||||
validate_stage_names,
|
||||
validate_stage_overrides,
|
||||
)
|
||||
from fastvideo.api.parser import (
|
||||
config_to_dict,
|
||||
load_config,
|
||||
@@ -31,6 +47,7 @@ from fastvideo.api.parser import (
|
||||
parse_config,
|
||||
)
|
||||
from fastvideo.api.results import GenerationResult
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
@@ -42,18 +59,26 @@ __all__ = [
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"GpuPoolConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"PromptEnhancerConfig",
|
||||
"PromptSafetyConfig",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"SamplingParam",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
"StreamingConfig",
|
||||
"WarmupConfig",
|
||||
"InferencePreset",
|
||||
"PresetStageSpec",
|
||||
"apply_overrides",
|
||||
"config_to_dict",
|
||||
"load_config",
|
||||
@@ -61,5 +86,12 @@ __all__ = [
|
||||
"load_run_config",
|
||||
"load_serve_config",
|
||||
"parse_cli_overrides",
|
||||
"get_all_preset_names",
|
||||
"get_preset",
|
||||
"get_presets_for_family",
|
||||
"parse_config",
|
||||
"register_preset",
|
||||
"validate_preset_selection",
|
||||
"validate_stage_names",
|
||||
"validate_stage_overrides",
|
||||
]
|
||||
|
||||
+211
-95
@@ -7,14 +7,17 @@ from dataclasses import fields, is_dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.overrides import apply_overrides, normalize_overrides
|
||||
from fastvideo.api.parser import config_to_dict, load_raw_config, parse_config
|
||||
from fastvideo.api.request_metadata import (
|
||||
EXPLICIT_REQUEST_ATTR,
|
||||
EXPLICIT_PATHS_ATTR,
|
||||
bind_generation_request_raw,
|
||||
refresh_generation_request_raw,
|
||||
get_explicit_paths,
|
||||
reset_tracking_roots,
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
@@ -22,8 +25,12 @@ from fastvideo.api.schema import (
|
||||
RequestRuntimeConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
from fastvideo.utils import shallow_asdict
|
||||
|
||||
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
|
||||
@@ -37,6 +44,10 @@ _LEGACY_REQUEST_ALIASES = {
|
||||
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
|
||||
"embedded_cfg_scale",
|
||||
})
|
||||
# torch.compile kwargs that map to first-class CompileConfig fields.
|
||||
_COMPILE_TYPED_KEYS = ("backend", "fullgraph", "mode", "dynamic")
|
||||
# LTX-2 refine flat kwargs (init + per-request) known to FastVideoArgs.
|
||||
_LTX2_REFINE_FLAT_KEYS = (refine_preset_override_fields() | refine_stage_override_fields())
|
||||
|
||||
|
||||
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
|
||||
@@ -50,7 +61,7 @@ def load_generator_config_from_file(
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> GeneratorConfig:
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = _normalize_overrides(overrides)
|
||||
normalized_overrides = normalize_overrides(overrides)
|
||||
|
||||
if _looks_like_run_or_serve_config(raw):
|
||||
if normalized_overrides:
|
||||
@@ -79,6 +90,8 @@ def legacy_from_pretrained_to_config(
|
||||
components: dict[str, Any] = {}
|
||||
quantization: dict[str, Any] = {}
|
||||
experimental: dict[str, Any] = {}
|
||||
preset_overrides: dict[str, Any] = {}
|
||||
preset_refine: dict[str, Any] = {}
|
||||
|
||||
for key, value in kwargs.items():
|
||||
if key == "revision":
|
||||
@@ -105,8 +118,33 @@ def legacy_from_pretrained_to_config(
|
||||
offload["pin_cpu_memory"] = value
|
||||
elif key == "enable_torch_compile":
|
||||
compile_config["enabled"] = value
|
||||
elif key == "enable_torch_compile_text_encoder":
|
||||
compile_config["text_encoder_enabled"] = value
|
||||
elif key == "torch_compile_kwargs":
|
||||
compile_config["kwargs"] = deepcopy(value)
|
||||
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
|
||||
for first_class in _COMPILE_TYPED_KEYS:
|
||||
if first_class in remaining:
|
||||
compile_config[first_class] = remaining.pop(first_class)
|
||||
if remaining:
|
||||
compile_config["extras"] = remaining
|
||||
elif key == "ltx2_vae_tiling":
|
||||
pipeline["vae_tiling"] = value
|
||||
elif key == "config_model_path":
|
||||
components["config_root"] = value
|
||||
elif key == "ltx2_refine_enabled":
|
||||
preset_refine["enabled"] = value
|
||||
elif key == "ltx2_refine_upsampler_path":
|
||||
# Empty string means "no upsampler"; keep typed None.
|
||||
components["upsampler_weights"] = value or None
|
||||
elif key == "ltx2_refine_lora_path":
|
||||
# Empty string means "no refine LoRA"; keep typed None.
|
||||
components["lora_path"] = value or None
|
||||
elif key == "ltx2_refine_add_noise":
|
||||
preset_refine["add_noise"] = value
|
||||
elif key == "ltx2_refine_num_inference_steps":
|
||||
preset_refine["num_inference_steps"] = value
|
||||
elif key == "ltx2_refine_guidance_scale":
|
||||
preset_refine["guidance_scale"] = value
|
||||
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
|
||||
engine[key] = value
|
||||
elif key == "override_text_encoder_quant":
|
||||
@@ -146,6 +184,10 @@ def legacy_from_pretrained_to_config(
|
||||
|
||||
if components:
|
||||
pipeline["components"] = components
|
||||
if preset_refine:
|
||||
preset_overrides["refine"] = preset_refine
|
||||
if preset_overrides:
|
||||
pipeline["preset_overrides"] = preset_overrides
|
||||
if experimental:
|
||||
pipeline["experimental"] = experimental
|
||||
if pipeline:
|
||||
@@ -157,16 +199,12 @@ def legacy_from_pretrained_to_config(
|
||||
def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, Any], ) -> FastVideoArgs:
|
||||
normalized = normalize_generator_config(config)
|
||||
unsupported = []
|
||||
if normalized.pipeline.profile is not None:
|
||||
unsupported.append("pipeline.profile")
|
||||
if normalized.pipeline.profile_version is not None:
|
||||
unsupported.append("pipeline.profile_version")
|
||||
if normalized.pipeline.components.config_root is not None:
|
||||
unsupported.append("pipeline.components.config_root")
|
||||
if normalized.pipeline.preset is not None:
|
||||
unsupported.append("pipeline.preset")
|
||||
if normalized.pipeline.preset_version is not None:
|
||||
unsupported.append("pipeline.preset_version")
|
||||
if normalized.pipeline.components.vae_weights is not None:
|
||||
unsupported.append("pipeline.components.vae_weights")
|
||||
if normalized.pipeline.components.upsampler_weights is not None:
|
||||
unsupported.append("pipeline.components.upsampler_weights")
|
||||
if unsupported:
|
||||
joined = ", ".join(unsupported)
|
||||
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
|
||||
@@ -190,13 +228,21 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
"vae_cpu_offload": engine.offload.vae,
|
||||
"pin_cpu_memory": engine.offload.pin_cpu_memory,
|
||||
"enable_torch_compile": engine.compile.enabled,
|
||||
"torch_compile_kwargs": deepcopy(engine.compile.kwargs),
|
||||
"torch_compile_kwargs": _compile_config_to_torch_kwargs(engine.compile),
|
||||
"enable_stage_verification": engine.enable_stage_verification,
|
||||
"use_fsdp_inference": engine.use_fsdp_inference,
|
||||
"disable_autocast": engine.disable_autocast,
|
||||
}
|
||||
if normalized.pipeline.workload_type is not None:
|
||||
kwargs["workload_type"] = normalized.pipeline.workload_type
|
||||
if normalized.pipeline.vae_tiling is not None:
|
||||
kwargs["ltx2_vae_tiling"] = normalized.pipeline.vae_tiling
|
||||
if engine.compile.text_encoder_enabled is not None:
|
||||
# ``FastVideoArgs.from_kwargs`` filters to declared fields, so
|
||||
# this is a no-op on the current legacy path. Emit anyway so the
|
||||
# realtime runtime (PR 7.6) — which reads from the kwargs dict
|
||||
# before FastVideoArgs filtering — can pick it up once wired.
|
||||
kwargs["enable_torch_compile_text_encoder"] = (engine.compile.text_encoder_enabled)
|
||||
|
||||
quantization = engine.quantization
|
||||
if quantization is not None and quantization.text_encoder_quant is not None:
|
||||
@@ -219,8 +265,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
kwargs["init_weights_from_safetensors"] = components.transformer_weights
|
||||
if components.transformer_2_weights is not None:
|
||||
kwargs["init_weights_from_safetensors_2"] = components.transformer_2_weights
|
||||
if components.config_root is not None:
|
||||
kwargs["config_model_path"] = components.config_root
|
||||
if components.upsampler_weights is not None:
|
||||
kwargs["ltx2_refine_upsampler_path"] = components.upsampler_weights
|
||||
|
||||
kwargs.update(deepcopy(normalized.pipeline.profile_overrides))
|
||||
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
|
||||
refine = preset_overrides.pop("refine", None)
|
||||
if isinstance(refine, Mapping):
|
||||
for key in _LTX2_REFINE_FLAT_KEYS:
|
||||
if key in refine:
|
||||
kwargs[f"ltx2_refine_{key}"] = refine[key]
|
||||
kwargs.update(preset_overrides)
|
||||
kwargs.update(deepcopy(normalized.pipeline.experimental))
|
||||
return FastVideoArgs.from_kwargs(**kwargs)
|
||||
|
||||
@@ -228,9 +284,9 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
def normalize_generation_request(request: GenerationRequest | Mapping[str, Any], ) -> GenerationRequest:
|
||||
normalized = (request if isinstance(request, GenerationRequest) else parse_config(GenerationRequest, request))
|
||||
|
||||
if hasattr(normalized, EXPLICIT_REQUEST_ATTR):
|
||||
refresh_generation_request_raw(normalized)
|
||||
else:
|
||||
if not hasattr(normalized, EXPLICIT_PATHS_ATTR):
|
||||
# Request wasn't bound through the parser (e.g. constructed
|
||||
# directly). Treat every currently-set field as explicit.
|
||||
bind_generation_request_raw(normalized, _serialize_generation_request(normalized))
|
||||
return normalized
|
||||
|
||||
@@ -270,16 +326,24 @@ def request_to_sampling_param(
|
||||
) -> SamplingParam:
|
||||
if request.plan is not None:
|
||||
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
|
||||
if request.state is not None:
|
||||
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
updates = _explicit_request_updates(request)
|
||||
if request.state is not None:
|
||||
_validate_continuation_state(request.state)
|
||||
sampling_param.continuation_state = request.state
|
||||
if request.output.return_state:
|
||||
sampling_param.return_continuation_state = True
|
||||
updates = explicit_request_updates(request)
|
||||
|
||||
for key, value in updates.items():
|
||||
if hasattr(sampling_param, key):
|
||||
setattr(sampling_param, key, deepcopy(value))
|
||||
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS or _is_supported_as_default_only(key, value):
|
||||
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
|
||||
continue
|
||||
elif value == _SCHEMA_DEFAULT_UPDATES.get(key, _MISSING):
|
||||
# Schema-default field that isn't on SamplingParam; tolerated
|
||||
# because direct GenerationRequest(...) construction has no
|
||||
# way to distinguish "user set" from "schema default".
|
||||
continue
|
||||
else:
|
||||
raise ValueError(f"Request field {key!r} is not supported by sampling params for {model_path}")
|
||||
@@ -296,10 +360,12 @@ def expand_request_prompt_batch(request: GenerationRequest, ) -> list[Generation
|
||||
requests: list[GenerationRequest] = []
|
||||
for index, prompt in enumerate(request.prompt):
|
||||
single_request = deepcopy(request)
|
||||
# deepcopy preserves the tracking-root cycle, but re-pin roots
|
||||
# defensively so that subsequent setattrs record on the copy.
|
||||
reset_tracking_roots(single_request)
|
||||
single_request.prompt = prompt
|
||||
_fan_out_batched_input_value(request, single_request, "image_path", index)
|
||||
_fan_out_batched_input_value(request, single_request, "video_path", index)
|
||||
_fan_out_explicit_request_metadata(request, single_request, index, prompt)
|
||||
requests.append(single_request)
|
||||
return requests
|
||||
|
||||
@@ -308,12 +374,23 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
|
||||
return isinstance(raw.get("generator"), Mapping)
|
||||
|
||||
|
||||
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
def _compile_config_to_torch_kwargs(compile_config: CompileConfig, ) -> dict[str, Any]:
|
||||
"""Flatten typed ``CompileConfig`` back to a ``torch_compile_kwargs``
|
||||
dict that the legacy ``FastVideoArgs`` path still expects.
|
||||
|
||||
Typed first-class fields (:attr:`backend`, :attr:`fullgraph`,
|
||||
:attr:`mode`, :attr:`dynamic`) are only emitted when the user set
|
||||
them explicitly (non-``None``). ``extras`` is merged on top for any
|
||||
uncommon kwargs.
|
||||
"""
|
||||
out: dict[str, Any] = {}
|
||||
for key in _COMPILE_TYPED_KEYS:
|
||||
value = getattr(compile_config, key)
|
||||
if value is not None:
|
||||
out[key] = value
|
||||
if compile_config.extras:
|
||||
out.update(deepcopy(compile_config.extras))
|
||||
return out
|
||||
|
||||
|
||||
def _sampling_param_to_request_raw(sampling_param: SamplingParam | None, ) -> dict[str, Any]:
|
||||
@@ -354,20 +431,83 @@ def _apply_request_field(
|
||||
|
||||
def request_to_pipeline_overrides(request: GenerationRequest) -> dict[str, Any]:
|
||||
overrides: dict[str, Any] = {}
|
||||
for key, value in _explicit_request_updates(request).items():
|
||||
for key, value in explicit_request_updates(request).items():
|
||||
if key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
|
||||
overrides[key] = deepcopy(value)
|
||||
return overrides
|
||||
|
||||
|
||||
def _explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
|
||||
raw = getattr(request, EXPLICIT_REQUEST_ATTR, None)
|
||||
if raw is None:
|
||||
raw = _serialize_generation_request(request)
|
||||
def explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
|
||||
"""Project a ``GenerationRequest`` down to *explicitly set* fields only.
|
||||
|
||||
Returns a flat kwargs dict suitable for merging into a generator call.
|
||||
The projection uses ``_fastvideo_explicit_paths`` (populated during
|
||||
``parse_config`` / raw binding) so schema defaults on the dataclass
|
||||
are **not** emitted — only paths the caller/operator actually wrote.
|
||||
|
||||
This is what makes ``ServeConfig.default_request`` work as an
|
||||
operator-pinned baseline rather than a full override: a YAML with just
|
||||
``sampling.seed: 42`` yields ``{"seed": 42}``, not the full sampling
|
||||
config with its 15 schema defaults.
|
||||
|
||||
Precondition: the request must carry ``_fastvideo_explicit_paths`` —
|
||||
populated by :func:`fastvideo.api.parser.parse_config` or
|
||||
:func:`fastvideo.api.compat.normalize_generation_request`. Calling on
|
||||
a raw ``GenerationRequest()`` asserts.
|
||||
"""
|
||||
assert hasattr(request,
|
||||
EXPLICIT_PATHS_ATTR), ("GenerationRequest reached explicit_request_updates without tracking; "
|
||||
"every entry point must route through normalize_generation_request "
|
||||
"or parse_config first")
|
||||
paths = get_explicit_paths(request)
|
||||
raw = _build_sparse_raw_from_paths(request, paths)
|
||||
return _extract_request_updates(raw)
|
||||
|
||||
|
||||
def _build_sparse_raw_from_paths(
|
||||
request: GenerationRequest,
|
||||
paths: frozenset[str],
|
||||
) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {}
|
||||
for path in paths:
|
||||
parts = path.split(".")
|
||||
value = _read_dotted_path(request, parts)
|
||||
if value is _MISSING:
|
||||
continue
|
||||
_set_dotted_path(result, parts, deepcopy(value))
|
||||
return result
|
||||
|
||||
|
||||
def _read_dotted_path(obj: Any, parts: list[str]) -> Any:
|
||||
for part in parts:
|
||||
if is_dataclass(obj) and not isinstance(obj, type):
|
||||
if not hasattr(obj, part):
|
||||
return _MISSING
|
||||
obj = getattr(obj, part)
|
||||
elif isinstance(obj, Mapping):
|
||||
if part not in obj:
|
||||
return _MISSING
|
||||
obj = obj[part]
|
||||
else:
|
||||
return _MISSING
|
||||
return obj
|
||||
|
||||
|
||||
def _set_dotted_path(
|
||||
target: dict[str, Any],
|
||||
parts: list[str],
|
||||
value: Any,
|
||||
) -> None:
|
||||
cursor = target
|
||||
for part in parts[:-1]:
|
||||
nxt = cursor.get(part)
|
||||
if not isinstance(nxt, dict):
|
||||
nxt = {}
|
||||
cursor[part] = nxt
|
||||
cursor = nxt
|
||||
cursor[parts[-1]] = value
|
||||
|
||||
|
||||
def _extract_request_updates(raw: Mapping[str, Any]) -> dict[str, Any]:
|
||||
updates: dict[str, Any] = {}
|
||||
if "negative_prompt" in raw:
|
||||
@@ -411,6 +551,40 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
|
||||
return deepcopy(config_to_dict(request))
|
||||
|
||||
|
||||
_SCHEMA_DEFAULT_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
|
||||
|
||||
_KNOWN_CONTINUATION_KINDS: set[str] = set()
|
||||
|
||||
|
||||
def register_continuation_kind(kind: str) -> None:
|
||||
"""Register a :class:`ContinuationState.kind` as recognized.
|
||||
|
||||
PR 7 wires the envelope through; per-kind payload deserializers live
|
||||
with each model family (e.g. ``fastvideo.pipelines.basic.ltx2.
|
||||
continuation.LTX2ContinuationState``). The registry lets the
|
||||
public-API compat layer validate the kind early, before the state
|
||||
reaches the pipeline.
|
||||
"""
|
||||
if not isinstance(kind, str) or not kind:
|
||||
raise ValueError("ContinuationState kind must be a non-empty string")
|
||||
_KNOWN_CONTINUATION_KINDS.add(kind)
|
||||
|
||||
|
||||
def _validate_continuation_state(state: ContinuationState) -> None:
|
||||
if not isinstance(state.kind, str) or not state.kind:
|
||||
raise ValueError("GenerationRequest.state.kind must be a non-empty string; got "
|
||||
f"{state.kind!r}")
|
||||
if not isinstance(state.payload, Mapping):
|
||||
raise ValueError(f"GenerationRequest.state.payload must be a mapping; got "
|
||||
f"{type(state.payload).__name__}")
|
||||
if state.kind not in _KNOWN_CONTINUATION_KINDS:
|
||||
known = sorted(_KNOWN_CONTINUATION_KINDS)
|
||||
raise ValueError(f"Unknown ContinuationState kind {state.kind!r}; registered "
|
||||
f"kinds: {known}. Import the model family that owns this kind "
|
||||
"(e.g. `import fastvideo.pipelines.basic.ltx2.continuation`) "
|
||||
"to register it, or drop the state field.")
|
||||
|
||||
|
||||
def _fan_out_batched_input_value(
|
||||
source_request: GenerationRequest,
|
||||
target_request: GenerationRequest,
|
||||
@@ -424,29 +598,6 @@ def _fan_out_batched_input_value(
|
||||
setattr(target_request.inputs, field_name, deepcopy(value[index]))
|
||||
|
||||
|
||||
def _fan_out_explicit_request_metadata(
|
||||
source_request: GenerationRequest,
|
||||
target_request: GenerationRequest,
|
||||
index: int,
|
||||
prompt: str,
|
||||
) -> None:
|
||||
raw = getattr(source_request, EXPLICIT_REQUEST_ATTR, None)
|
||||
if raw is None:
|
||||
return
|
||||
|
||||
raw = deepcopy(raw)
|
||||
raw["prompt"] = prompt
|
||||
inputs = raw.get("inputs")
|
||||
if isinstance(inputs, dict):
|
||||
for field_name in ("image_path", "video_path"):
|
||||
value = inputs.get(field_name)
|
||||
if isinstance(value, list):
|
||||
_validate_batched_input_length(source_request.prompt, value, field_name)
|
||||
inputs[field_name] = deepcopy(value[index])
|
||||
|
||||
setattr(target_request, EXPLICIT_REQUEST_ATTR, raw)
|
||||
|
||||
|
||||
def _validate_batched_input_length(
|
||||
prompts: str | list[str] | None,
|
||||
values: list[Any],
|
||||
@@ -458,50 +609,15 @@ def _validate_batched_input_length(
|
||||
raise ValueError(f"GenerationRequest.inputs.{field_name} must have the same length as request.prompt")
|
||||
|
||||
|
||||
def _is_supported_as_default_only(key: str, value: Any) -> bool:
|
||||
default_value = _DEFAULT_REQUEST_UPDATES.get(key, _MISSING)
|
||||
return default_value is not _MISSING and _values_equal(value, default_value)
|
||||
|
||||
|
||||
def _collect_non_default_fields(
|
||||
value: Any,
|
||||
default: Any,
|
||||
) -> dict[str, Any]:
|
||||
if not (is_dataclass(value) and is_dataclass(default)):
|
||||
return {}
|
||||
|
||||
result: dict[str, Any] = {}
|
||||
for field in fields(value):
|
||||
current = getattr(value, field.name)
|
||||
default_value = getattr(default, field.name)
|
||||
if is_dataclass(current) and is_dataclass(default_value):
|
||||
nested = _collect_non_default_fields(current, default_value)
|
||||
if nested:
|
||||
result[field.name] = nested
|
||||
continue
|
||||
if not _values_equal(current, default_value):
|
||||
result[field.name] = deepcopy(current)
|
||||
return result
|
||||
|
||||
|
||||
def _values_equal(left: Any, right: Any) -> bool:
|
||||
if left is right:
|
||||
return True
|
||||
try:
|
||||
return bool(left == right)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
_DEFAULT_REQUEST_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
|
||||
|
||||
__all__ = [
|
||||
"explicit_request_updates",
|
||||
"generator_config_to_fastvideo_args",
|
||||
"legacy_from_pretrained_to_config",
|
||||
"legacy_generate_call_to_request",
|
||||
"load_generator_config_from_file",
|
||||
"normalize_generation_request",
|
||||
"normalize_generator_config",
|
||||
"register_continuation_kind",
|
||||
"request_to_pipeline_overrides",
|
||||
"request_to_sampling_param",
|
||||
]
|
||||
|
||||
@@ -45,6 +45,15 @@ def apply_overrides(config: Mapping[str, Any], overrides: Mapping[str, Any]) ->
|
||||
return merged
|
||||
|
||||
|
||||
def normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
"""Normalize a CLI list or mapping of overrides into a flat dict."""
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
|
||||
|
||||
def _apply_single_override(config: dict[str, Any], dotted_key: str, value: Any) -> None:
|
||||
parts = dotted_key.split(".")
|
||||
if not all(parts):
|
||||
@@ -98,4 +107,4 @@ def _normalize_override_key(key: str) -> str:
|
||||
return key.replace("-", "_")
|
||||
|
||||
|
||||
__all__ = ["apply_overrides", "parse_cli_overrides"]
|
||||
__all__ = ["apply_overrides", "normalize_overrides", "parse_cli_overrides"]
|
||||
|
||||
+2
-10
@@ -11,7 +11,7 @@ from typing import Any, Literal, TypeVar, Union, get_args, get_origin, get_type_
|
||||
import yaml
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.overrides import apply_overrides, normalize_overrides
|
||||
from fastvideo.api.request_metadata import (
|
||||
bind_generation_request_raw,
|
||||
bind_run_config_raw,
|
||||
@@ -64,7 +64,7 @@ def load_config(
|
||||
) -> T:
|
||||
"""Load a typed config object from YAML or JSON."""
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = _normalize_overrides(overrides)
|
||||
normalized_overrides = normalize_overrides(overrides)
|
||||
if normalized_overrides:
|
||||
raw = apply_overrides(raw, normalized_overrides)
|
||||
return parse_config(config_type, raw)
|
||||
@@ -108,14 +108,6 @@ def _load_raw_mapping(handle: Any, config_path: Path) -> Any:
|
||||
raise ValueError(f"Unsupported config file format: {config_path}")
|
||||
|
||||
|
||||
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
|
||||
|
||||
class _SchemaParser:
|
||||
|
||||
def parse_dataclass(
|
||||
|
||||
@@ -0,0 +1,261 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pipeline preset registry.
|
||||
|
||||
A *preset* is a named inference preset for a model family. It bundles:
|
||||
|
||||
* ``defaults`` — sampling values applied when the user does not
|
||||
override them (consumed at runtime via ``SamplingParam.from_pretrained``);
|
||||
* ``stage_schemas`` — **validation-only** metadata describing which
|
||||
user-facing stage names (``"denoise"``, ``"sr"``) the preset recognises
|
||||
and which ``stage_overrides`` keys each stage accepts.
|
||||
|
||||
The ``stage_schemas`` tuple does **not** drive pipeline execution. The
|
||||
concrete execution DAG (text encoding, denoising, VAE decoding, …) is
|
||||
hard-coded per-pipeline in ``create_pipeline_stages()``. Schemas exist
|
||||
purely so that ``PipelineSelection.preset`` and
|
||||
``GenerationRequest.stage_overrides`` can be type-checked up front
|
||||
without touching the pipeline.
|
||||
|
||||
Preset base types and the registry API live here (public API surface).
|
||||
Preset *instances* are defined in pipeline-local ``presets.py`` files
|
||||
(e.g. ``fastvideo/pipelines/basic/wan/presets.py``) and registered
|
||||
explicitly from :func:`_register_presets` in ``fastvideo/registry.py``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Types
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PresetStageSpec:
|
||||
"""A user-facing stage name within a preset, used only to validate
|
||||
``stage_overrides`` keys. Not read by pipeline execution — the real
|
||||
execution DAG lives in each pipeline's ``create_pipeline_stages()``.
|
||||
"""
|
||||
|
||||
name: str
|
||||
"""Short user-facing name, e.g. ``"denoise"``, ``"sr"``."""
|
||||
|
||||
kind: str
|
||||
"""Semantic kind, e.g. ``"denoising"``, ``"super_resolution"``."""
|
||||
|
||||
description: str = ""
|
||||
|
||||
allowed_overrides: frozenset[str] = field(default_factory=frozenset)
|
||||
"""Keys that may appear in ``stage_overrides[name]``."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class InferencePreset:
|
||||
"""A named inference preset for a model family."""
|
||||
|
||||
name: str
|
||||
"""Preset name, e.g. ``"wan_t2v_1_3b"``."""
|
||||
|
||||
version: int
|
||||
"""Preset schema version; bump on breaking schema changes."""
|
||||
|
||||
model_family: str
|
||||
"""Model family key, e.g. ``"wan"``, ``"ltx2"``."""
|
||||
|
||||
description: str = ""
|
||||
|
||||
workload_type: str | None = None
|
||||
"""Optional workload hint: ``"t2v"``, ``"i2v"``, etc."""
|
||||
|
||||
stage_schemas: tuple[PresetStageSpec, ...] = ()
|
||||
"""User-facing stage names for ``stage_overrides`` validation.
|
||||
|
||||
Validation-only: this tuple is consumed by
|
||||
:func:`validate_stage_overrides` and is **not** used to drive
|
||||
pipeline execution. Omit or leave empty if the preset exposes no
|
||||
per-stage override surface.
|
||||
"""
|
||||
|
||||
defaults: dict[str, Any] = field(default_factory=dict)
|
||||
"""Preset-level default sampling/runtime values."""
|
||||
|
||||
stage_defaults: dict[str, dict[str, Any]] = field(default_factory=dict)
|
||||
"""Per-stage default overrides, keyed by stage name."""
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Registry
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
# Keyed by (model_family, name, version).
|
||||
_PRESET_REGISTRY: dict[tuple[str, str, int], InferencePreset] = {}
|
||||
|
||||
|
||||
def register_preset(preset: InferencePreset) -> None:
|
||||
"""Register a preset definition.
|
||||
|
||||
Raises :class:`ValueError` on duplicate
|
||||
``(model_family, name, version)`` keys.
|
||||
"""
|
||||
key = (preset.model_family, preset.name, preset.version)
|
||||
if key in _PRESET_REGISTRY:
|
||||
raise ValueError(f"Duplicate preset registration: "
|
||||
f"model_family={key[0]!r}, name={key[1]!r}, "
|
||||
f"version={key[2]!r}")
|
||||
_PRESET_REGISTRY[key] = preset
|
||||
|
||||
|
||||
def get_preset(
|
||||
name: str,
|
||||
model_family: str,
|
||||
version: int | None = None,
|
||||
) -> InferencePreset:
|
||||
"""Look up a registered preset.
|
||||
|
||||
When *version* is ``None`` the highest registered version for the
|
||||
given *(model_family, name)* pair is returned.
|
||||
|
||||
Raises :class:`~fastvideo.api.errors.ConfigValidationError` when the
|
||||
preset cannot be found.
|
||||
"""
|
||||
if version is not None:
|
||||
key = (model_family, name, version)
|
||||
preset = _PRESET_REGISTRY.get(key)
|
||||
if preset is not None:
|
||||
return preset
|
||||
raise ConfigValidationError(
|
||||
"pipeline.preset",
|
||||
f"unknown preset {name!r} version {version!r} "
|
||||
f"for model family {model_family!r}; "
|
||||
f"registered: {_format_registered(model_family)}",
|
||||
)
|
||||
|
||||
# Find the highest version for (model_family, name).
|
||||
candidates = [prof for (fam, n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family and n == name]
|
||||
if not candidates:
|
||||
raise ConfigValidationError(
|
||||
"pipeline.preset",
|
||||
f"unknown preset {name!r} for model family "
|
||||
f"{model_family!r}; "
|
||||
f"registered: {_format_registered(model_family)}",
|
||||
)
|
||||
return max(candidates, key=lambda p: p.version)
|
||||
|
||||
|
||||
def get_presets_for_family(model_family: str, ) -> list[InferencePreset]:
|
||||
"""Return all presets registered for *model_family*."""
|
||||
return [prof for (fam, _n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family]
|
||||
|
||||
|
||||
def get_all_preset_names() -> list[str]:
|
||||
"""Return the sorted list of all registered preset names."""
|
||||
return sorted({prof.name for prof in _PRESET_REGISTRY.values()})
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Validation helpers
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
def validate_stage_names(
|
||||
preset: InferencePreset,
|
||||
stage_overrides: Mapping[str, Any],
|
||||
) -> None:
|
||||
"""Check that *stage_overrides* keys are valid stage names.
|
||||
|
||||
Raises :class:`~fastvideo.api.errors.ConfigValidationError` with a
|
||||
path-qualified message for unknown stage names.
|
||||
"""
|
||||
valid_names = {stage.name for stage in preset.stage_schemas}
|
||||
for stage_name in stage_overrides:
|
||||
if stage_name not in valid_names:
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}",
|
||||
f"unknown stage for preset {preset.name!r}; "
|
||||
f"valid stages: {sorted(valid_names)}",
|
||||
)
|
||||
|
||||
|
||||
def validate_stage_overrides(
|
||||
preset: InferencePreset,
|
||||
stage_overrides: Mapping[str, Any],
|
||||
) -> None:
|
||||
"""Validate stage override keys against the preset.
|
||||
|
||||
Calls :func:`validate_stage_names` first, then checks that each
|
||||
override key is in the stage's ``allowed_overrides``.
|
||||
"""
|
||||
validate_stage_names(preset, stage_overrides)
|
||||
stages_by_name = {stage.name: stage for stage in preset.stage_schemas}
|
||||
for stage_name, overrides in stage_overrides.items():
|
||||
if not isinstance(overrides, Mapping):
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}",
|
||||
"must be a mapping",
|
||||
)
|
||||
stage_spec = stages_by_name[stage_name]
|
||||
if not stage_spec.allowed_overrides:
|
||||
if overrides:
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}",
|
||||
f"stage {stage_name!r} does not accept "
|
||||
f"overrides",
|
||||
)
|
||||
continue
|
||||
for key in overrides:
|
||||
if key not in stage_spec.allowed_overrides:
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}.{key}",
|
||||
f"not an allowed override for stage "
|
||||
f"{stage_name!r}; allowed: "
|
||||
f"{sorted(stage_spec.allowed_overrides)}",
|
||||
)
|
||||
|
||||
|
||||
def validate_preset_selection(
|
||||
preset_name: str | None,
|
||||
model_family: str,
|
||||
*,
|
||||
preset_version: int | None = None,
|
||||
stage_overrides: Mapping[str, Any] | None = None,
|
||||
) -> InferencePreset | None:
|
||||
"""Resolve and validate a preset selection end-to-end.
|
||||
|
||||
Returns the resolved :class:`InferencePreset`, or ``None`` if
|
||||
*preset_name* is ``None`` (no preset requested).
|
||||
"""
|
||||
if preset_name is None:
|
||||
return None
|
||||
preset = get_preset(preset_name, model_family, version=preset_version)
|
||||
if stage_overrides:
|
||||
validate_stage_overrides(preset, stage_overrides)
|
||||
return preset
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
def _format_registered(model_family: str) -> str:
|
||||
names = sorted({prof.name for (fam, _n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family})
|
||||
if not names:
|
||||
return "(none)"
|
||||
return ", ".join(repr(n) for n in names)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InferencePreset",
|
||||
"PresetStageSpec",
|
||||
"get_all_preset_names",
|
||||
"get_preset",
|
||||
"get_presets_for_family",
|
||||
"register_preset",
|
||||
"validate_preset_selection",
|
||||
"validate_stage_names",
|
||||
"validate_stage_overrides",
|
||||
]
|
||||
@@ -1,24 +1,25 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Track which GenerationRequest fields the user explicitly provided.
|
||||
|
||||
This module solves a specific problem: when translating a GenerationRequest into
|
||||
a legacy SamplingParam, we need to distinguish user-provided values (which
|
||||
should override model defaults) from schema defaults (which should NOT override
|
||||
model defaults).
|
||||
When translating a GenerationRequest into a legacy SamplingParam we must
|
||||
distinguish user-provided values (which should override model defaults)
|
||||
from schema defaults (which should NOT override model defaults).
|
||||
|
||||
The approach:
|
||||
1. At bind time, store the original raw dict and a baseline snapshot.
|
||||
2. Patch __setattr__ on tracked dataclass types to record dirty field paths.
|
||||
3. At access time, do a lazy 3-way merge: raw + baseline + current state,
|
||||
with dirty paths forcing inclusion even when current == baseline.
|
||||
The mechanism: a single ``_fastvideo_explicit_paths`` set stored on the
|
||||
root ``GenerationRequest``. It holds dotted leaf paths (e.g.
|
||||
``"sampling.guidance_scale"``) the user has touched, either via raw
|
||||
config at bind time or via attribute assignment at runtime. A patched
|
||||
``__setattr__`` on the request dataclass types records assignments into
|
||||
this set.
|
||||
|
||||
The set holds leaf paths only. Nested dataclass or mapping assignments
|
||||
are flattened to their leaves at record time.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from collections.abc import Callable, Mapping
|
||||
import dataclasses
|
||||
from typing import Any, cast
|
||||
from collections.abc import Callable
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
@@ -33,12 +34,12 @@ from fastvideo.api.schema import (
|
||||
ServeConfig,
|
||||
)
|
||||
|
||||
EXPLICIT_REQUEST_ATTR = "_fastvideo_explicit_request"
|
||||
ORIGINAL_REQUEST_STATE_ATTR = "_fastvideo_original_request_state"
|
||||
EXPLICIT_PATHS_ATTR = "_fastvideo_explicit_paths"
|
||||
|
||||
_TRACKING_ROOT_ATTR = "_fastvideo_request_tracking_root"
|
||||
_TRACKING_PATH_ATTR = "_fastvideo_request_tracking_path"
|
||||
_TRACKING_PATCHED_ATTR = "_fastvideo_request_tracking_patched"
|
||||
_DIRTY_PATHS_ATTR = "_fastvideo_dirty_paths"
|
||||
|
||||
_TRACKED_REQUEST_TYPES = (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
@@ -55,14 +56,20 @@ def bind_generation_request_raw(
|
||||
request: GenerationRequest,
|
||||
raw: Mapping[str, Any] | None,
|
||||
) -> GenerationRequest:
|
||||
"""Install explicit-path tracking on *request*.
|
||||
|
||||
*raw* is the parsed config dict (YAML/JSON/kwargs); every leaf key
|
||||
in it becomes an explicit path. Subsequent attribute assignments on
|
||||
*request* or its nested dataclasses are recorded automatically via a
|
||||
patched ``__setattr__``.
|
||||
"""
|
||||
_ensure_request_tracking()
|
||||
# Disable dirty tracking during bind so tree walk doesn't record paths.
|
||||
object.__setattr__(request, _DIRTY_PATHS_ATTR, None)
|
||||
object.__setattr__(request, EXPLICIT_REQUEST_ATTR, deepcopy(dict(raw or {})))
|
||||
object.__setattr__(request, ORIGINAL_REQUEST_STATE_ATTR, _serialize_config(request))
|
||||
# Disable recording while we walk the tree to install roots.
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, None)
|
||||
_set_tracking_roots(request, request, "")
|
||||
# Enable dirty tracking.
|
||||
object.__setattr__(request, _DIRTY_PATHS_ATTR, set())
|
||||
paths: set[str] = set()
|
||||
_record_value_paths(raw or {}, "", paths)
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, paths)
|
||||
return request
|
||||
|
||||
|
||||
@@ -73,6 +80,8 @@ def bind_run_config_raw(
|
||||
request_raw = raw.get("request")
|
||||
if isinstance(request_raw, Mapping):
|
||||
bind_generation_request_raw(config.request, request_raw)
|
||||
else:
|
||||
bind_generation_request_raw(config.request, {})
|
||||
return config
|
||||
|
||||
|
||||
@@ -83,87 +92,72 @@ def bind_serve_config_raw(
|
||||
default_request_raw = raw.get("default_request")
|
||||
if isinstance(default_request_raw, Mapping):
|
||||
bind_generation_request_raw(config.default_request, default_request_raw)
|
||||
elif "default_request" not in raw:
|
||||
else:
|
||||
bind_generation_request_raw(config.default_request, {})
|
||||
return config
|
||||
|
||||
|
||||
def refresh_generation_request_raw(request: GenerationRequest, ) -> dict[str, Any] | None:
|
||||
raw = getattr(request, EXPLICIT_REQUEST_ATTR, None)
|
||||
baseline = getattr(request, ORIGINAL_REQUEST_STATE_ATTR, None)
|
||||
if not isinstance(raw, Mapping) or not isinstance(baseline, Mapping):
|
||||
return None
|
||||
def get_explicit_paths(request: GenerationRequest) -> frozenset[str]:
|
||||
"""Return a snapshot of the explicit paths set on *request*."""
|
||||
paths = getattr(request, EXPLICIT_PATHS_ATTR, None)
|
||||
if isinstance(paths, set | frozenset):
|
||||
return frozenset(paths)
|
||||
return frozenset()
|
||||
|
||||
dirty = getattr(request, _DIRTY_PATHS_ATTR, None) or frozenset()
|
||||
current = _serialize_config(request)
|
||||
merged = deepcopy(dict(raw))
|
||||
_merge_request_mutations(merged, dict(baseline), current, dirty)
|
||||
|
||||
object.__setattr__(request, EXPLICIT_REQUEST_ATTR, merged)
|
||||
object.__setattr__(request, ORIGINAL_REQUEST_STATE_ATTR, current)
|
||||
object.__setattr__(request, _DIRTY_PATHS_ATTR, set())
|
||||
return merged
|
||||
def reset_tracking_roots(request: GenerationRequest) -> None:
|
||||
"""Re-install tracking roots after a deepcopy or manual clone.
|
||||
|
||||
The paths set itself deepcopies correctly; we only need to repoint
|
||||
the tracking root on nested dataclasses at the new root.
|
||||
"""
|
||||
_ensure_request_tracking()
|
||||
_set_tracking_roots(request, request, "")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3-way merge: raw + baseline + current, with dirty-path forcing
|
||||
# Path recording
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_MISSING = object()
|
||||
|
||||
|
||||
def _merge_request_mutations(
|
||||
merged: dict[str, Any],
|
||||
baseline: Mapping[str, Any],
|
||||
current: Mapping[str, Any],
|
||||
dirty: frozenset[str] | set[str],
|
||||
path_prefix: str = "",
|
||||
force_dirty: bool = False,
|
||||
def _record_value_paths(
|
||||
value: Any,
|
||||
prefix: str,
|
||||
out: set[str],
|
||||
) -> None:
|
||||
# Remove keys that were deleted from the current state.
|
||||
for key in set(merged) | set(baseline):
|
||||
if key not in current:
|
||||
merged.pop(key, None)
|
||||
"""Add every leaf path under *value* to *out*.
|
||||
|
||||
for key in current:
|
||||
current_path = f"{path_prefix}.{key}" if path_prefix else key
|
||||
current_value = current[key]
|
||||
baseline_value = baseline.get(key, _MISSING)
|
||||
merged_value = merged.get(key, _MISSING)
|
||||
|
||||
# If this exact path was dirtied (e.g. whole section replaced),
|
||||
# propagate to all children.
|
||||
child_force = force_dirty or current_path in dirty
|
||||
|
||||
# Recurse into nested mappings.
|
||||
if isinstance(current_value, Mapping) and isinstance(baseline_value, Mapping):
|
||||
nested = (deepcopy(dict(merged_value)) if isinstance(merged_value, Mapping) else {})
|
||||
_merge_request_mutations(
|
||||
nested,
|
||||
baseline_value,
|
||||
current_value,
|
||||
dirty,
|
||||
current_path,
|
||||
child_force,
|
||||
)
|
||||
if nested:
|
||||
merged[key] = nested
|
||||
else:
|
||||
merged.pop(key, None)
|
||||
continue
|
||||
|
||||
# A field is explicitly set if:
|
||||
# - it's new (not in baseline),
|
||||
# - it changed from baseline,
|
||||
# - its path was touched by __setattr__ (dirty), or
|
||||
# - an ancestor path was dirty (whole section replaced).
|
||||
is_dirty = child_force or current_path in dirty
|
||||
if baseline_value is _MISSING or current_value != baseline_value or is_dirty:
|
||||
merged[key] = deepcopy(current_value)
|
||||
A leaf is any terminal value (non-dataclass, non-mapping, or empty
|
||||
mapping/dataclass). ``prefix`` is the dotted path at which *value*
|
||||
sits. When called with an empty ``prefix`` (the root), leaves are
|
||||
recorded at their own key.
|
||||
"""
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
dc_fields = dataclasses.fields(value)
|
||||
if not dc_fields:
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
return
|
||||
for field in dc_fields:
|
||||
child = getattr(value, field.name)
|
||||
path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
_record_value_paths(child, path, out)
|
||||
return
|
||||
if isinstance(value, Mapping):
|
||||
if not value:
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
return
|
||||
for key, child in value.items():
|
||||
path = f"{prefix}.{key}" if prefix else key
|
||||
_record_value_paths(child, path, out)
|
||||
return
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# __setattr__ patching for dirty-path recording
|
||||
# __setattr__ patching
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@@ -187,16 +181,23 @@ def _patch_tracking_setattr(config_type: type[Any]) -> None:
|
||||
original_setattr(self, name, value)
|
||||
return
|
||||
|
||||
root = getattr(self, _TRACKING_ROOT_ATTR, None)
|
||||
if root is not None:
|
||||
dirty = getattr(root, _DIRTY_PATHS_ATTR, None)
|
||||
if isinstance(dirty, set):
|
||||
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
|
||||
path = f"{prefix}.{name}" if prefix else name
|
||||
dirty.add(path)
|
||||
|
||||
original_setattr(self, name, value)
|
||||
|
||||
root = getattr(self, _TRACKING_ROOT_ATTR, None)
|
||||
if root is None:
|
||||
return
|
||||
paths = getattr(root, EXPLICIT_PATHS_ATTR, None)
|
||||
if not isinstance(paths, set):
|
||||
return
|
||||
|
||||
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
|
||||
path = f"{prefix}.{name}" if prefix else name
|
||||
# Wholesale dataclass replacement: install roots on the new
|
||||
# instance so its future mutations are tracked too.
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
_set_tracking_roots(root, value, path)
|
||||
_record_value_paths(value, path, paths)
|
||||
|
||||
type.__setattr__(config_type, "__setattr__", _tracking_setattr)
|
||||
setattr(config_type, _TRACKING_PATCHED_ATTR, True)
|
||||
|
||||
@@ -222,26 +223,11 @@ def _set_tracking_roots(
|
||||
_set_tracking_roots(root, child, child_path)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialization helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _serialize_config(config: Any) -> Any:
|
||||
if dataclasses.is_dataclass(config) and not isinstance(config, type):
|
||||
return {field.name: _serialize_config(getattr(config, field.name)) for field in dataclasses.fields(config)}
|
||||
if isinstance(config, list):
|
||||
return [_serialize_config(item) for item in config]
|
||||
if isinstance(config, dict):
|
||||
return {key: _serialize_config(value) for key, value in config.items()}
|
||||
return deepcopy(config)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EXPLICIT_REQUEST_ATTR",
|
||||
"ORIGINAL_REQUEST_STATE_ATTR",
|
||||
"EXPLICIT_PATHS_ATTR",
|
||||
"bind_generation_request_raw",
|
||||
"bind_run_config_raw",
|
||||
"bind_serve_config_raw",
|
||||
"refresh_generation_request_raw",
|
||||
"get_explicit_paths",
|
||||
"reset_tracking_roots",
|
||||
]
|
||||
|
||||
@@ -15,6 +15,7 @@ class GenerationResult:
|
||||
samples: Any | None = None
|
||||
frames: Any | None = None
|
||||
audio: Any | None = None
|
||||
audio_sample_rate: int | None = None
|
||||
size: tuple[int, int, int] | None = None
|
||||
generation_time: float | None = None
|
||||
logging_info: Any | None = None
|
||||
@@ -44,6 +45,7 @@ class GenerationResult:
|
||||
"samples",
|
||||
"frames",
|
||||
"audio",
|
||||
"audio_sample_rate",
|
||||
"size",
|
||||
"generation_time",
|
||||
"logging_info",
|
||||
@@ -62,6 +64,7 @@ class GenerationResult:
|
||||
samples=result.get("samples"),
|
||||
frames=result.get("frames"),
|
||||
audio=result.get("audio"),
|
||||
audio_sample_rate=result.get("audio_sample_rate"),
|
||||
size=result.get("size"),
|
||||
generation_time=result.get("generation_time"),
|
||||
logging_info=result.get("logging_info"),
|
||||
@@ -80,6 +83,7 @@ class GenerationResult:
|
||||
"samples": self.samples,
|
||||
"frames": self.frames,
|
||||
"audio": self.audio,
|
||||
"audio_sample_rate": self.audio_sample_rate,
|
||||
"size": self.size,
|
||||
"generation_time": self.generation_time,
|
||||
"logging_info": self.logging_info,
|
||||
|
||||
@@ -1,10 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import StoreBoolean
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -30,6 +36,16 @@ class SamplingParam:
|
||||
|
||||
# Camera control inputs (HYWorld)
|
||||
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
|
||||
# Camera/action control inputs (GameCraft)
|
||||
camera_states: Any | None = None # Plücker coordinates [B, T_video, 6, H, W]
|
||||
camera_trajectory: str | None = None
|
||||
action_list: list[str] | None = None
|
||||
action_speed_list: list[float] | None = None
|
||||
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
|
||||
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
|
||||
|
||||
# Camera control inputs (LingBotWorld)
|
||||
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
@@ -68,6 +84,7 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
@@ -80,6 +97,54 @@ class SamplingParam:
|
||||
movement_distance: float | None = None
|
||||
camera_rotation: str | None = None
|
||||
|
||||
# LTX-2 multi-modal CFG and STG.
|
||||
# cfg_scale defaults are 1.0 (CFG off) so ``ForwardBatch.__post_init__``
|
||||
# doesn't force ``do_classifier_free_guidance`` on non-LTX-2 models that
|
||||
# never override these fields. LTX-2 presets that need text-CFG on set
|
||||
# them in their ``defaults`` dict (e.g. ``ltx2_base``).
|
||||
ltx2_cfg_scale_video: float = 1.0
|
||||
ltx2_cfg_scale_audio: float = 1.0
|
||||
ltx2_modality_scale_video: float = 3.0
|
||||
ltx2_modality_scale_audio: float = 3.0
|
||||
ltx2_rescale_scale: float = 0.7
|
||||
ltx2_stg_scale_video: float = 1.0
|
||||
ltx2_stg_scale_audio: float = 1.0
|
||||
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
|
||||
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
|
||||
|
||||
# Stable Audio (T2A): clip start/end in seconds. Honored by
|
||||
# `StableAudioConditioningStage` + `StableAudioDecodingStage`. Other
|
||||
# families ignore them.
|
||||
audio_start_in_s: float | None = None
|
||||
audio_end_in_s: float | None = None
|
||||
|
||||
# Stable Audio audio-to-audio (variation):
|
||||
# `init_audio` -- a path or `[B, C, samples]` waveform at the model
|
||||
# sample rate; the pipeline encodes it via the VAE
|
||||
# and uses it as the starting latent.
|
||||
# `init_audio_strength` -- 0..1, higher = closer to the reference
|
||||
# (matches the convention of Stability's
|
||||
# commercial Stable Audio 2.0 UI). 1.0 ~=
|
||||
# VAE round-trip, 0.0 ~= plain T2A.
|
||||
# `init_noise_level` -- legacy raw `sigma_max` override (0.3..500,
|
||||
# higher = more freedom). Kept for callers
|
||||
# that already use it; prefer `init_audio_strength`.
|
||||
init_audio: Any = None
|
||||
init_audio_strength: float | None = None
|
||||
init_noise_level: float | None = None
|
||||
|
||||
# Stable Audio inpainting (RePaint-style): `inpaint_audio` is the
|
||||
# reference clip, `inpaint_mask` is a [samples] tensor in {0, 1} where
|
||||
# 1 means *keep the reference* and 0 means *regenerate*.
|
||||
inpaint_audio: Any = None
|
||||
inpaint_mask: Any = None
|
||||
|
||||
# Continuation state carried across streaming/multi-segment calls.
|
||||
continuation_state: ContinuationState | None = None
|
||||
# When True, the pipeline returns a ContinuationState on the result so
|
||||
# the caller can resume from the generated segment.
|
||||
return_continuation_state: bool = False
|
||||
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
@@ -94,26 +159,58 @@ class SamplingParam:
|
||||
raise ValueError("prompt_path must be a txt file")
|
||||
|
||||
def update(self, source_dict: dict[str, Any]) -> None:
|
||||
valid_fields = {f.name for f in fields(self)}
|
||||
for key, value in source_dict.items():
|
||||
if hasattr(self, key):
|
||||
if key in valid_fields:
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
logger.exception("%s has no attribute %s", type(self).__name__, key)
|
||||
logger.error("%s has no field %s", type(self).__name__, key)
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.registry import get_sampling_param_cls_for_name
|
||||
sampling_cls = get_sampling_param_cls_for_name(model_path)
|
||||
if sampling_cls is not None:
|
||||
sampling_param: SamplingParam = sampling_cls()
|
||||
else:
|
||||
logger.warning("Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
model_path)
|
||||
sampling_param = cls()
|
||||
def from_pretrained(cls, model_path: str) -> SamplingParam:
|
||||
sampling_param = cls._from_preset(model_path)
|
||||
if sampling_param is not None:
|
||||
return sampling_param
|
||||
|
||||
return sampling_param
|
||||
logger.warning(
|
||||
"Couldn't find a preset for %s."
|
||||
" Using the default sampling param.",
|
||||
model_path,
|
||||
)
|
||||
return cls()
|
||||
|
||||
@classmethod
|
||||
def _from_preset(
|
||||
cls,
|
||||
model_path: str,
|
||||
) -> SamplingParam | None:
|
||||
"""Build a SamplingParam from preset defaults.
|
||||
|
||||
Returns ``None`` when no preset is configured for
|
||||
*model_path*, letting the caller fall back to the legacy
|
||||
subclass lookup.
|
||||
"""
|
||||
from fastvideo.registry import get_preset_selection
|
||||
|
||||
try:
|
||||
preset_name, model_family = get_preset_selection(model_path)
|
||||
except (ValueError, RuntimeError):
|
||||
return None
|
||||
if preset_name is None or model_family is None:
|
||||
return None
|
||||
|
||||
from fastvideo.api.presets import get_preset
|
||||
|
||||
preset = get_preset(preset_name, model_family)
|
||||
sp = cls()
|
||||
valid_fields = {f.name for f in fields(cls)}
|
||||
for key, value in preset.defaults.items():
|
||||
if key in valid_fields:
|
||||
setattr(sp, key, copy.deepcopy(value))
|
||||
sp.__post_init__()
|
||||
return sp
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: Any) -> Any:
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user