Compare commits
64
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c2e89f22d3 | ||
|
|
e47a3c5aad | ||
|
|
d40fbfc534 | ||
|
|
126a52ce32 | ||
|
|
419e1c68f1 | ||
|
|
938bc3c972 | ||
|
|
381a7aae66 | ||
|
|
7fa4fb50bb | ||
|
|
de3cd6aab2 | ||
|
|
a64e5e62ab | ||
|
|
f4200cc3f3 | ||
|
|
363230b372 | ||
|
|
1dcdacc35a | ||
|
|
0a7a4a21f6 | ||
|
|
c30d731992 | ||
|
|
f8eaff4292 | ||
|
|
76e7048b3d | ||
|
|
b1386a7a78 | ||
|
|
664d6b3c23 | ||
|
|
0688c5131a | ||
|
|
0b605d0a41 | ||
|
|
8f5ebe2aeb | ||
|
|
84803076a0 | ||
|
|
5cc337fdb6 | ||
|
|
2b8e5a56a8 | ||
|
|
e117167eb9 | ||
|
|
4f9af56b78 | ||
|
|
33efe7a673 | ||
|
|
b07cc3f9d4 | ||
|
|
29eb4109cc | ||
|
|
d07a7691fe | ||
|
|
960485519f | ||
|
|
412c95b1a0 | ||
|
|
f6275f8005 | ||
|
|
7b872cc41e | ||
|
|
37418946c8 | ||
|
|
95fd29e0cb | ||
|
|
e17cd2633c | ||
|
|
e0dc5f2b0c | ||
|
|
70ee5d230c | ||
|
|
24ced500f5 | ||
|
|
4ddcdf541f | ||
|
|
0e3529869c | ||
|
|
e1e0d91c00 | ||
|
|
145a3f166b | ||
|
|
88a5a933ab | ||
|
|
c591d6d2a6 | ||
|
|
65dff806a8 | ||
|
|
b85f0f4c2a | ||
|
|
76c62d7a00 | ||
|
|
f6e65ff668 | ||
|
|
c220aa8000 | ||
|
|
4713fc17ed | ||
|
|
5789955bbe | ||
|
|
2ad84a3b78 | ||
|
|
12d699cd78 | ||
|
|
34f14ded21 | ||
|
|
71d1ab411f | ||
|
|
805e487773 | ||
|
|
8803b4547e | ||
|
|
3b3806b3f6 | ||
|
|
38d962e89d | ||
|
|
3966a365d0 | ||
|
|
d73fd14af0 |
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,4 @@
|
||||
{"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"}
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
---
|
||||
name: seed-ssim-references
|
||||
description: Seed HF reference videos for a single newly-added SSIM test. Runs the test on Modal L40S, downloads the generated mp4s via `modal volume get`, pauses for the user to eyeball quality, 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 Videos
|
||||
|
||||
## Purpose
|
||||
|
||||
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
|
||||
reference videos exist on the HF dataset (`FastVideo/ssim-reference-videos`).
|
||||
This skill:
|
||||
|
||||
1. Runs the test on Modal's L40S pool to generate the videos.
|
||||
2. Downloads them to the local repo via `modal volume get`.
|
||||
3. Pauses so the user can eyeball the mp4s and confirm quality.
|
||||
4. 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 mp4 without crashing. 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
|
||||
|
||||
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.
|
||||
|
||||
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:
|
||||
|
||||
```bash
|
||||
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
|
||||
```
|
||||
|
||||
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
|
||||
|
||||
Print the list of downloaded mp4s and their paths, then stop. Tell the user:
|
||||
|
||||
> "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."
|
||||
|
||||
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 mp4s. Loop over each `<model_id>` extracted
|
||||
in step 1:
|
||||
|
||||
```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`
|
||||
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: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
|
||||
### 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.
|
||||
- **Modal run fails before generation.** No mp4s 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 videos (most likely the test
|
||||
crashed before writing, or `REQUIRED_GPUS` exceeded the partition capacity
|
||||
— see Modal logs).
|
||||
- **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 mp4s 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.
|
||||
|
||||
## 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 (SSIM drifts across SKUs).
|
||||
- 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`.
|
||||
|
||||
## 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.
|
||||
- `fastvideo/tests/ssim/README.md` — reference layout, HF repo conventions.
|
||||
- `fastvideo/tests/ssim/inference_similarity_utils.py` —
|
||||
`run_text_to_video_similarity_test` + `_build_init_kwargs`: what each test
|
||||
config passes to `VideoGenerator.from_pretrained`.
|
||||
|
||||
## 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. |
|
||||
@@ -4,21 +4,24 @@ description: How to develop, validate, and register a new evaluation metric
|
||||
|
||||
# Evaluation Development SOP
|
||||
|
||||
Standard procedure for adding new video quality evaluation metrics to the
|
||||
FastVideo agent toolkit.
|
||||
Standard procedure for adding new video quality evaluation metrics to
|
||||
the FastVideo agent toolkit.
|
||||
|
||||
## When to Use
|
||||
## When to use
|
||||
|
||||
- You need a metric that doesn't exist in `.agents/memory/evaluation-registry/README.md`.
|
||||
- You need a metric that does not exist in
|
||||
`.agents/memory/evaluation-registry/README.md`.
|
||||
- An existing metric needs significant changes to its methodology.
|
||||
- You're exploring a new evaluation approach.
|
||||
- You are exploring a new evaluation approach.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Research
|
||||
|
||||
- Search `.agents/memory/related-work/` for existing evaluation approaches.
|
||||
- Check the `evaluation_registry.md` for current metrics and their limitations.
|
||||
- Search `.agents/memory/related-work/` for existing evaluation
|
||||
approaches.
|
||||
- Check `.agents/memory/evaluation-registry/README.md` for current
|
||||
metrics and their limitations.
|
||||
- Review literature: FVD, CLIP-Score, human preference, etc.
|
||||
|
||||
### 2. Prototype
|
||||
@@ -29,21 +32,25 @@ FastVideo agent toolkit.
|
||||
|
||||
### 3. Validate
|
||||
|
||||
- **Known-good test**: Metric should score high on reference-quality videos.
|
||||
- **Known-bad test**: Metric should score low on degraded/unrelated videos.
|
||||
- **Sensitivity test**: Small quality differences should produce meaningful
|
||||
score differences.
|
||||
- **Known-good test**: metric should score high on reference-quality
|
||||
videos.
|
||||
- **Known-bad test**: metric should score low on degraded or unrelated
|
||||
videos.
|
||||
- **Sensitivity test**: small quality differences should produce
|
||||
meaningful score differences.
|
||||
- Document thresholds and their justification.
|
||||
|
||||
### 4. Register
|
||||
|
||||
Update `.agents/memory/evaluation-registry/README.md`:
|
||||
|
||||
- Add the metric with status `Active`.
|
||||
- Document location, thresholds, and trust level.
|
||||
|
||||
### 5. Integrate
|
||||
|
||||
Update `.agents/skills/evaluate-video-quality.md`:
|
||||
Update `.agents/skills/evaluate-video-quality/SKILL.md`:
|
||||
|
||||
- Add the new metric as a section.
|
||||
- Include code examples and interpretation guide.
|
||||
|
||||
@@ -52,3 +59,35 @@ Update `.agents/skills/evaluate-video-quality.md`:
|
||||
- Move the exploration log content into the skill.
|
||||
- Clean up the exploration file or mark it as `promoted`.
|
||||
- If anything went wrong during development, create a lesson.
|
||||
|
||||
## Where the metrics live
|
||||
|
||||
The eval suite is `fastvideo/eval/`. New metrics register themselves
|
||||
via `@register("<group>.<name>")` and are auto-discovered when
|
||||
`fastvideo.eval.metrics` is imported.
|
||||
|
||||
- **Native metrics** (SSIM, PSNR, LPIPS, optical flow, VLM): add a
|
||||
file under the appropriate group dir
|
||||
(`fastvideo/eval/metrics/common/`, `optical_flow/`, `videoscore2/`,
|
||||
`physics_iq/`).
|
||||
- **Metrics that wrap upstream research code**: follow the vbench
|
||||
pattern in `fastvideo/eval/metrics/vbench/`. The contract is:
|
||||
- Upstream lives as a git submodule under
|
||||
`fastvideo/third_party/eval/<bench>/`, pinned to a SHA in repo-root
|
||||
`.gitmodules`.
|
||||
- The metric package's `__init__.py` inserts the submodule path on
|
||||
`sys.path` and installs runtime compat shims (attribute-level
|
||||
monkey-patches) for any modern-dep drift. Do not modify upstream
|
||||
files on disk, and do not ship a `setup.sh`.
|
||||
- See `fastvideo/eval/README.md` for the worked vbench example.
|
||||
- Full porting guide:
|
||||
[`docs/contributing/eval-metrics.md`](../../docs/contributing/eval-metrics.md).
|
||||
|
||||
## Out of scope of the initial eval port
|
||||
|
||||
The following land in follow-up PRs:
|
||||
|
||||
- **MIND** metrics (depends on a separate `vipe` submodule).
|
||||
- **VBench-2.0** sibling package.
|
||||
- Native conversion of **FVD** under `fastvideo/eval/metrics/fvd/`.
|
||||
- The training-time `EvalCallback`.
|
||||
|
||||
+186
-2
@@ -9,11 +9,183 @@ notify:
|
||||
- github_commit_status:
|
||||
context: "full-suite-passed"
|
||||
if: build.env("TEST_SCOPE") == "full"
|
||||
- github_commit_status:
|
||||
context: "direct-test-completed"
|
||||
if: build.env("TEST_SCOPE") == "direct"
|
||||
|
||||
steps:
|
||||
# ============================================================
|
||||
- label: ":dart: Direct Test (${TEST_TYPE})"
|
||||
if: build.env("TEST_SCOPE") == "direct"
|
||||
# Direct test: triggered by /test <name> slash command.
|
||||
# Labels match fastcheck/full-suite counterparts so the GitHub
|
||||
# check status overwrites the original failed check.
|
||||
# Only ONE step executes per build (gated by TEST_TYPE).
|
||||
# ============================================================
|
||||
|
||||
# --- Fastcheck-scope direct tests ---
|
||||
- label: ":microscope: Encoder Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "encoder"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: VAE Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "vae"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: Transformer Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "transformer"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: Kernel Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "kernel_tests"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":microscope: Unit Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "unit_test"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
# --- Full-suite-scope direct tests ---
|
||||
- label: ":bar_chart: SSIM Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "ssim"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Inference Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "inference_lora"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Distillation DMD Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "distillation_dmd"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Self-Forcing Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "self_forcing"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training_lora"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests VSA"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training_vsa"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Inference Tests VMoBA"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "inference_vmoba"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Performance Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "performance"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: API Server Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "api_server"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
@@ -135,6 +307,10 @@ steps:
|
||||
label: ":bar_chart: SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
@@ -195,6 +371,10 @@ steps:
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
@@ -207,6 +387,10 @@ steps:
|
||||
label: ":test_tube: Training Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=training_vsa
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
|
||||
+7
-12
@@ -4,8 +4,10 @@ merge_protections:
|
||||
- base = main
|
||||
success_conditions:
|
||||
- "title~=(?i)^\\[(feat|feature|bugfix|fix|refactor|perf|ci|doc|docs|misc|chore|kernel|new.?model)\\]"
|
||||
- "#approved-reviews-by>=1"
|
||||
- check-success~=pre-commit
|
||||
- check-success=fastcheck-passed
|
||||
- check-success=full-suite-passed
|
||||
|
||||
pull_request_rules:
|
||||
|
||||
@@ -103,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
|
||||
@@ -272,24 +274,15 @@ pull_request_rules:
|
||||
merge:
|
||||
method: squash
|
||||
|
||||
- name: auto-rebase when ready and Full Suite passed
|
||||
- name: auto-update when ready
|
||||
conditions:
|
||||
- label=ready
|
||||
- "#approved-reviews-by>=1"
|
||||
- check-success=full-suite-passed
|
||||
- -conflict
|
||||
- -closed
|
||||
- -draft
|
||||
actions:
|
||||
rebase: {}
|
||||
|
||||
- name: remove ready label on Full Suite failure
|
||||
conditions:
|
||||
- label=ready
|
||||
- check-failure=full-suite-passed
|
||||
actions:
|
||||
label:
|
||||
remove: [ready]
|
||||
update: {}
|
||||
|
||||
# ============================================================
|
||||
# PR title format help
|
||||
@@ -319,3 +312,5 @@ pull_request_rules:
|
||||
|
||||
Please update your PR title and the merge protection check will pass automatically.
|
||||
|
||||
merge_protections_settings:
|
||||
reporting_method: check-runs
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
name: Aggregate Test Status
|
||||
|
||||
on:
|
||||
status:
|
||||
|
||||
permissions:
|
||||
statuses: write
|
||||
|
||||
jobs:
|
||||
aggregate:
|
||||
if: >-
|
||||
github.event.context == 'direct-test-completed'
|
||||
&& github.event.state == 'success'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check and update aggregate status
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const sha = context.payload.sha;
|
||||
|
||||
const { data } = await github.rest.repos.getCombinedStatusForRef({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
ref: sha,
|
||||
per_page: 100,
|
||||
});
|
||||
|
||||
const bkStatuses = data.statuses.filter(
|
||||
s => s.context.startsWith('buildkite/ci/')
|
||||
);
|
||||
|
||||
const FASTCHECK_PREFIX = 'buildkite/ci/microscope-';
|
||||
const FULL_SUITE_PREFIXES = [
|
||||
'buildkite/ci/test-tube-',
|
||||
'buildkite/ci/bar-chart-',
|
||||
];
|
||||
|
||||
const fastcheck = bkStatuses.filter(
|
||||
s => s.context.startsWith(FASTCHECK_PREFIX)
|
||||
);
|
||||
const fullSuite = bkStatuses.filter(
|
||||
s => FULL_SUITE_PREFIXES.some(p => s.context.startsWith(p))
|
||||
);
|
||||
|
||||
if (
|
||||
fastcheck.length > 0
|
||||
&& fastcheck.every(s => s.state === 'success')
|
||||
) {
|
||||
core.info(
|
||||
`All ${fastcheck.length} fastcheck tests passed — updating fastcheck-passed`
|
||||
);
|
||||
await github.rest.repos.createCommitStatus({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
sha,
|
||||
state: 'success',
|
||||
context: 'fastcheck-passed',
|
||||
description:
|
||||
`All ${fastcheck.length} fastcheck tests passed`,
|
||||
});
|
||||
}
|
||||
|
||||
if (
|
||||
fullSuite.length > 0
|
||||
&& fullSuite.every(s => s.state === 'success')
|
||||
) {
|
||||
core.info(
|
||||
`All ${fullSuite.length} full suite tests passed — updating full-suite-passed`
|
||||
);
|
||||
await github.rest.repos.createCommitStatus({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
sha,
|
||||
state: 'success',
|
||||
context: 'full-suite-passed',
|
||||
description:
|
||||
`All ${fullSuite.length} full suite tests passed`,
|
||||
});
|
||||
}
|
||||
@@ -4,10 +4,11 @@ on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
|
||||
concurrency:
|
||||
group: pre-commit-${{ github.ref }}
|
||||
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
|
||||
inputs:
|
||||
ref:
|
||||
description: 'Git ref to checkout (defaults to github.ref)'
|
||||
required: false
|
||||
type: string
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -18,6 +19,8 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
@@ -33,6 +33,7 @@ jobs:
|
||||
core.setOutput('has_write', String(hasWrite));
|
||||
|
||||
- name: Add ready label and react
|
||||
id: label
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
@@ -40,7 +41,6 @@ jobs:
|
||||
const owner = context.repo.owner;
|
||||
const repo = context.repo.repo;
|
||||
const prNumber = context.payload.issue.number;
|
||||
// Remove ready first to allow re-trigger (labeled event fires on add, not if already present)
|
||||
try { await github.rest.issues.removeLabel({ owner, repo, issue_number: prNumber, name: 'ready' }); } catch {}
|
||||
await github.rest.issues.addLabels({ owner, repo, issue_number: prNumber, labels: ['ready'] });
|
||||
await github.rest.reactions.createForIssueComment({
|
||||
@@ -48,6 +48,44 @@ jobs:
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
const { data: pr } = await github.rest.pulls.get({ owner, repo, pull_number: prNumber });
|
||||
core.setOutput('pr_sha', pr.head.sha);
|
||||
core.setOutput('pr_branch', pr.head.ref);
|
||||
core.setOutput('pr_number', String(prNumber));
|
||||
|
||||
- name: Trigger Full Suite
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_SHA: ${{ steps.label.outputs.pr_sha }}
|
||||
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
|
||||
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
curl -sS --fail-with-body -X POST \
|
||||
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds" \
|
||||
-H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
-H "Content-Type: application/json" \
|
||||
--data-raw "$(jq -n \
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
branch: $branch,
|
||||
message: $message,
|
||||
ignore_pipeline_branch_filters: true,
|
||||
pull_request_id: $pr_id,
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
}
|
||||
}')"
|
||||
|
||||
parse-command:
|
||||
if: >-
|
||||
github.event.issue.pull_request != null
|
||||
@@ -143,12 +181,26 @@ jobs:
|
||||
core.setOutput('sha', pr.head.sha);
|
||||
core.setOutput('branch', pr.head.ref);
|
||||
|
||||
- name: React to comment
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
await github.rest.reactions.createForIssueComment({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
|
||||
pre-commit:
|
||||
needs: parse-command
|
||||
if: >-
|
||||
needs.parse-command.outputs.has_write == 'true'
|
||||
&& needs.parse-command.outputs.test_scope == 'precommit'
|
||||
uses: ./.github/workflows/ci-precommit.yml
|
||||
with:
|
||||
ref: refs/pull/${{ github.event.issue.number }}/merge
|
||||
|
||||
post-precommit-status:
|
||||
needs: [parse-command, pre-commit]
|
||||
@@ -178,17 +230,6 @@ jobs:
|
||||
&& needs.parse-command.outputs.test_type != ''
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: React to comment
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
await github.rest.reactions.createForIssueComment({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
|
||||
- name: Trigger Buildkite
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
name: Trigger Full Suite
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
pull_request_target:
|
||||
types: [labeled, synchronize]
|
||||
|
||||
permissions:
|
||||
@@ -10,7 +10,7 @@ permissions:
|
||||
|
||||
concurrency:
|
||||
group: full-suite-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
trigger:
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
# Find running builds for this branch with TEST_SCOPE=full and cancel them
|
||||
builds=$(curl -sS -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds?branch=${PR_BRANCH}&state=running,scheduled" \
|
||||
| jq -r '.[] | select(.env.TEST_SCOPE == "full") | .number')
|
||||
| jq -r '.[] | select(try (.env.TEST_SCOPE == "full") catch false) | .number')
|
||||
for build_num in $builds; do
|
||||
echo "Cancelling Buildkite build #$build_num"
|
||||
curl -sS -X PUT -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
|
||||
@@ -85,6 +85,7 @@ docs/distillation/examples/
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
# Next.js / Node artifacts under ui/: see ui/.gitignore
|
||||
|
||||
.claude/
|
||||
.codex/
|
||||
|
||||
@@ -4,3 +4,6 @@
|
||||
[submodule "fastvideo-kernel/include/cutlass"]
|
||||
path = fastvideo-kernel/include/cutlass
|
||||
url = https://github.com/NVIDIA/cutlass.git
|
||||
[submodule "fastvideo/third_party/eval/vbench"]
|
||||
path = fastvideo/third_party/eval/vbench
|
||||
url = https://github.com/Vchitect/VBench.git
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
WRN 2026-03-26T13:46:33.469 ?.19646 server_start:193: Failed to start server: operation not permitted: /var/folders/z_/h_6myyk14d1b7z87z3vy4mjh0000gn/T/nvim.dsynkd/iSe0el/nvim.19646.0
|
||||
@@ -0,0 +1 @@
|
||||
3.12
|
||||
@@ -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
|
||||
|
||||
@@ -24,7 +24,7 @@ PR push
|
||||
Runs on the PR branch directly
|
||||
│
|
||||
pass ──► Mergify auto-squash-merges to main, branch deleted
|
||||
fail ──► Mergify removes 'ready' label; fix and /merge again
|
||||
fail ──► fix the regression, push, and /merge again
|
||||
```
|
||||
|
||||
---
|
||||
@@ -102,8 +102,8 @@ failing test's output.
|
||||
| Performance Tests | `performance` | 30 min |
|
||||
| API Server Tests | `api_server` | 30 min |
|
||||
|
||||
A Full Suite failure removes the `ready` label automatically. A Mergify comment links to
|
||||
the Buildkite build. Fix the regression, push, and comment `/merge` again.
|
||||
If a Full Suite test fails, check the Buildkite build log for the failing step's output.
|
||||
Fix the regression, push, and comment `/merge` again to re-trigger.
|
||||
|
||||
---
|
||||
|
||||
@@ -129,8 +129,8 @@ Suite passing directly on the PR branch.
|
||||
- No merge conflicts
|
||||
5. If all conditions pass, Mergify squash-merges to `main` automatically. The branch is
|
||||
deleted after merge.
|
||||
6. If the Full Suite fails, Mergify removes the `ready` label and posts a comment linking to
|
||||
the Buildkite build. The developer fixes the issue, pushes, and comments `/merge` again.
|
||||
6. If the Full Suite fails, the developer fixes the issue, pushes, and comments `/merge`
|
||||
again to re-trigger.
|
||||
|
||||
**Merge conditions summary:**
|
||||
|
||||
@@ -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/` |
|
||||
@@ -279,6 +279,30 @@ Triggers a specific Buildkite test or suite on the current PR branch.
|
||||
| `/test api` | API server integration tests | `api_server` |
|
||||
| `/test full` | Entire Full Suite | all (with `TEST_SCOPE=full`) |
|
||||
| `/test fastcheck` | Entire Fastcheck suite | fastcheck (with `TEST_SCOPE=fastcheck`) |
|
||||
| `/test pre-commit` | Pre-commit checks on PR code | — (runs `ci-precommit.yml` via `workflow_call`) |
|
||||
|
||||
**Re-running failed tests:** When you use `/test <name>` to re-run a specific failed test,
|
||||
the resulting Buildkite check uses the same name as the original (e.g., `/test encoder`
|
||||
creates `buildkite/ci/microscope-encoder-tests`). This overwrites the failed check status.
|
||||
Once all tests in a tier pass, the aggregate status (`fastcheck-passed` or
|
||||
`full-suite-passed`) is automatically updated to `success` by the `ci-aggregate-status.yml`
|
||||
workflow.
|
||||
|
||||
**How aggregate status refresh works:**
|
||||
|
||||
1. `/test <name>` triggers a Buildkite build with `TEST_SCOPE=direct`. The test step uses
|
||||
the same label as its fastcheck/full-suite counterpart, so the resulting GitHub check
|
||||
overwrites the original.
|
||||
2. When the build completes, Buildkite's `notify` posts a `direct-test-completed` commit
|
||||
status. This is the only signal that triggers the aggregate workflow — intermediate step
|
||||
status updates do not trigger it.
|
||||
3. `ci-aggregate-status.yml` fires, calls `getCombinedStatusForRef` to fetch the latest
|
||||
status for every context on that commit (each context returns only its most recent
|
||||
state), groups them by prefix (`microscope-*` → fastcheck, `test-tube-*`/`bar-chart-*`
|
||||
→ full suite), and posts `fastcheck-passed: success` or `full-suite-passed: success` if
|
||||
all entries in the group are `success`.
|
||||
4. Tests that were never triggered (skipped by monorepo-diff) have no status entry and do
|
||||
not block the aggregate.
|
||||
|
||||
---
|
||||
|
||||
@@ -296,6 +320,7 @@ Protected branches (`main`, `master`, `release/*`) are never deleted.
|
||||
| `ci-precommit.yml` | Every push / PR against `main` | Runs pre-commit hooks (yapf, ruff, mypy, codespell, pymarkdown, actionlint, check-filenames) |
|
||||
| `ci-trigger-full-suite.yml` | `ready` label added to a PR | Calls Buildkite API to run Full Suite on the PR branch |
|
||||
| `ci-slash-commands.yml` | PR comment starting with `/merge` or `/test` | Handles slash commands; adds `ready` label or triggers Buildkite |
|
||||
| `ci-aggregate-status.yml` | Any Buildkite commit status update | Checks if all tests in a tier passed; updates `fastcheck-passed` or `full-suite-passed` |
|
||||
| `community-issue-labeler.yml` | Issue opened or edited | Auto-labels issues by keyword matching against title and body |
|
||||
| `community-welcome.yml` | First contribution | Posts a welcome comment for first-time contributors |
|
||||
| `community-stale.yml` | Scheduled | Marks and closes stale issues and PRs |
|
||||
|
||||
@@ -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)
|
||||
@@ -319,7 +319,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:
|
||||
|
||||
@@ -474,7 +475,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`
|
||||
|
||||
@@ -0,0 +1,559 @@
|
||||
# Porting Eval Metrics into `fastvideo.eval`
|
||||
|
||||
This guide is for contributors adding new evaluation metrics to
|
||||
FastVideo's eval suite. To run the existing metrics, see
|
||||
[`fastvideo/eval/README.md`](../../fastvideo/eval/README.md).
|
||||
|
||||
## When to use this guide
|
||||
|
||||
Use this guide when you are:
|
||||
|
||||
- Adding a new metric (native or wrapping a third-party library).
|
||||
- Porting a benchmark (e.g. VBench, MIND, EvalCrafter) whose Python
|
||||
code needs to be importable from a pinned upstream.
|
||||
- Adding a new metric group (audio, vlm, etc.).
|
||||
|
||||
## TL;DR
|
||||
|
||||
Metrics are auto-discovered from
|
||||
`fastvideo/eval/metrics/<group>/<name>/metric.py`. Each declares itself
|
||||
with `@register("<group>.<name>")` and subclasses `BaseMetric`. Three
|
||||
recipes:
|
||||
|
||||
1. **Native metric** (pure-PyTorch, no submodule). Drop a file,
|
||||
declare deps, implement `compute(sample)`.
|
||||
2. **Library-wrapped metric** (CLIP, torch.hub, transformers, pyiqa).
|
||||
Same as above, plus route the library's cache through
|
||||
`get_cache_dir()` if it has a `download_root=` / `cache_dir=`
|
||||
kwarg.
|
||||
3. **Upstream-submodule-wrapped metric** (vbench-style). Pin upstream
|
||||
as a git submodule under `fastvideo/third_party/eval/<bench>/`. The
|
||||
adapter `__init__.py` does the `sys.path` insert and any runtime
|
||||
compat shims for modern dep versions. Patches live as Python in
|
||||
that file rather than as on-disk patches to the submodule.
|
||||
|
||||
The full recipes are below.
|
||||
|
||||
---
|
||||
|
||||
## 0) Layout and auto-discovery
|
||||
|
||||
```
|
||||
fastvideo/eval/metrics/
|
||||
├── base.py # BaseMetric + lifecycle contract
|
||||
├── common/ # group: SSIM, PSNR, LPIPS
|
||||
├── optical_flow/ # group: gt_optical_flow, synthetic_optical_flow
|
||||
├── vlm/ # group: VideoScore-2
|
||||
├── physics_iq/ # group + sub-metrics
|
||||
└── vbench/ # group: 16 sub-metrics
|
||||
├── __init__.py # sys.path bootstrap + runtime compat shims
|
||||
├── _grit_helper.py # shared upstream-touching helpers
|
||||
└── <sub_metric>/metric.py
|
||||
```
|
||||
|
||||
Auto-discovery (`fastvideo/eval/metrics/__init__.py`) walks each group
|
||||
dir and imports every `metric.py` it finds, which fires the
|
||||
`@register` decorators. Names starting with `_` are skipped. Use that
|
||||
prefix for shared helpers or vendored code that should not register
|
||||
itself.
|
||||
|
||||
---
|
||||
|
||||
## 1) The `BaseMetric` contract
|
||||
|
||||
Every metric subclasses `fastvideo.eval.metrics.base.BaseMetric` and
|
||||
declares:
|
||||
|
||||
```python
|
||||
class YourMetric(BaseMetric):
|
||||
name: str = "common.your_metric" # must match @register
|
||||
requires_reference: bool = True # needs sample["reference"]
|
||||
higher_is_better: bool = True # for ranking / aggregates
|
||||
dependencies: list[str] = [] # importable module names;
|
||||
# registry surfaces a clean
|
||||
# ImportError if missing
|
||||
needs_gpu: bool = False
|
||||
backbone: str | None = None # e.g. "clip_vit_l14"
|
||||
```
|
||||
|
||||
You must implement:
|
||||
|
||||
```python
|
||||
def compute(self, sample: dict) -> list[MetricResult]:
|
||||
"""sample['video'] is (1, T, C, H, W). Return a one-element list.
|
||||
|
||||
The leading 1 is preserved for forward-compat with batched eval;
|
||||
today :class:`EvalWorker` always invokes metrics with B=1.
|
||||
"""
|
||||
```
|
||||
|
||||
You may override:
|
||||
|
||||
- `setup(self) -> None`. Eager model loading. Called once by
|
||||
`create_evaluator`. Idempotent (re-entrant). Use the `if self._model
|
||||
is not None: return` pattern.
|
||||
- `to(self, device)`. Move the metric and its submodels to `device`.
|
||||
|
||||
If a required input is missing (e.g. an fps-aware metric called
|
||||
without `fps`), return `self._skip(sample, reason)` instead of
|
||||
raising.
|
||||
|
||||
---
|
||||
|
||||
## 2) Recipe A: native metric (no external deps)
|
||||
|
||||
Smallest case. Pixel math, simple closed-form.
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/common/your_metric/metric.py
|
||||
from __future__ import annotations
|
||||
import torch
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult
|
||||
|
||||
|
||||
@register("common.your_metric")
|
||||
class YourMetric(BaseMetric):
|
||||
name = "common.your_metric"
|
||||
requires_reference = True
|
||||
higher_is_better = True
|
||||
needs_gpu = False
|
||||
dependencies: list[str] = [] # nothing extra
|
||||
|
||||
def compute(self, sample: dict) -> list[MetricResult]:
|
||||
gen, ref = sample["video"], sample["reference"] # (B,T,C,H,W) each
|
||||
per_video = ((gen - ref) ** 2).mean(dim=(1, 2, 3, 4)).sqrt()
|
||||
return [
|
||||
MetricResult(name=self.name, score=float(s), details={})
|
||||
for s in per_video
|
||||
]
|
||||
```
|
||||
|
||||
That is the whole recipe. Drop the file and the registry picks it up.
|
||||
|
||||
---
|
||||
|
||||
## 3) Recipe B: library-wrapped metric (CLIP, torch.hub, transformers, pyiqa)
|
||||
|
||||
If your metric loads a backbone from a Python package, route the
|
||||
library at the eval cache so users get one knob (`FASTVIDEO_EVAL_CACHE`)
|
||||
to redirect everything.
|
||||
|
||||
### Cache routing rules
|
||||
|
||||
| Library | How to route | Location after redirect |
|
||||
|---|---|---|
|
||||
| `clip.load("ViT-X")` | pass `download_root=str(get_cache_dir() / "clip")` | `${FASTVIDEO_EVAL_CACHE}/clip/` |
|
||||
| `torch.hub.load(...)` | nothing; `TORCH_HOME` is redirected at `fastvideo.eval` import time | `${FASTVIDEO_EVAL_CACHE}/torch/hub/` |
|
||||
| `transformers.from_pretrained(...)` | nothing; leave HF's default cache (`~/.cache/huggingface/hub/`) so users dedupe with other ML projects | `~/.cache/huggingface/hub/` |
|
||||
| `huggingface_hub.snapshot_download` / `hf_hub_download` | use `ensure_checkpoint(...)` (it wraps these with filelock) | same as above |
|
||||
| `pyiqa.create_metric(...)` | no env var or kwarg honored; document in metric docstring | pyiqa-internal |
|
||||
| `lpips`, `ptlflow` | torch.hub-based, auto-redirected | `${FASTVIDEO_EVAL_CACHE}/torch/hub/` |
|
||||
| Raw URL (no HF Hub) | use `ensure_checkpoint(name, source="https://...")` | `${FASTVIDEO_EVAL_CACHE}/models/<name>` |
|
||||
| Dataset asset (raw video/mask/image) auto-fetched from a public bucket | download into `get_cache_dir() / "datasets" / "<bench>"`, mirroring upstream's relative layout. Vendor any small manifest (CSV/JSON ≤1 MB) under the metric folder so the dataset can be used without external setup. | `${FASTVIDEO_EVAL_CACHE}/datasets/<bench>/` |
|
||||
|
||||
### Dataset assets: vendor the manifest, auto-fetch the rest
|
||||
|
||||
If your metric ships with its own paired-reference dataset (Physics-IQ
|
||||
is the canonical example), follow this layout:
|
||||
|
||||
- **Manifest** (CSV/JSON ≤1 MB): vendor it under
|
||||
`fastvideo/eval/metrics/<bench>/_vendored/<manifest>.<ext>`, with a
|
||||
sibling `_vendored/LICENSE` recording attribution and provenance.
|
||||
The `_vendored/` subdir is the project-wide convention for
|
||||
upstream-provenance files: it is auto-skipped by metric discovery
|
||||
(the `_` prefix) and by codespell (one `*/_vendored/*` glob in
|
||||
`[tool.codespell].skip`), so dropping in a new vendored file
|
||||
requires no further config. Read the manifest from the dataset
|
||||
module via a `Path(__file__)`-relative resolver. Mirror
|
||||
`_VENDORED_DESCRIPTIONS_CSV` in
|
||||
`fastvideo/eval/datasets/physics_iq.py`.
|
||||
- **Heavy assets** (videos, masks, images): do not vendor. Auto-fetch
|
||||
on first miss into `get_cache_dir() / "datasets" / "<bench>"`,
|
||||
mirroring upstream's relative directory layout one-for-one so a
|
||||
pre-downloaded mirror at any path works as a drop-in
|
||||
`dataset_root=`. Use atomic `.part` then final-rename to be safe
|
||||
under concurrent SLURM ranks.
|
||||
- **Bucket override**: expose `FASTVIDEO_<BENCH>_BUCKET_URL` so users
|
||||
with internal mirrors can redirect.
|
||||
- **Opt-out**: accept `auto_download: bool = True` in the dataset
|
||||
constructor; on `False`, raise `FileNotFoundError` instead of
|
||||
fetching. This covers air-gapped runs and CI.
|
||||
|
||||
The end-state is `get_dataset("<bench>")` with no kwargs.
|
||||
|
||||
### Example: CLIP backbone + LAION head
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/your_group/your_metric/metric.py
|
||||
from __future__ import annotations
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult
|
||||
|
||||
|
||||
@register("your_group.your_metric")
|
||||
class YourMetric(BaseMetric):
|
||||
name = "your_group.your_metric"
|
||||
requires_reference = False
|
||||
needs_gpu = True
|
||||
dependencies = ["clip"] # "openai-clip" PyPI; importable as `clip`
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._clip = None
|
||||
self._head = None
|
||||
|
||||
def setup(self) -> None:
|
||||
if self._clip is not None:
|
||||
return
|
||||
import clip
|
||||
from fastvideo.eval.models import ensure_checkpoint, get_cache_dir
|
||||
|
||||
# Backbone: route CLIP's cache through our root.
|
||||
self._clip, _ = clip.load(
|
||||
"ViT-L/14",
|
||||
device=self.device,
|
||||
download_root=str(get_cache_dir() / "clip"),
|
||||
)
|
||||
self._clip.eval()
|
||||
|
||||
# URL-fetched head: ensure_checkpoint downloads to
|
||||
# ${FASTVIDEO_EVAL_CACHE}/models/ with filelock + atomic rename.
|
||||
ckpt = ensure_checkpoint(
|
||||
"your_head.pth",
|
||||
source="https://example.com/path/to/your_head.pth",
|
||||
)
|
||||
self._head = nn.Linear(768, 1)
|
||||
self._head.load_state_dict(
|
||||
torch.load(ckpt, map_location="cpu", weights_only=True)
|
||||
)
|
||||
self._head.to(self.device).eval()
|
||||
|
||||
def to(self, device):
|
||||
super().to(device)
|
||||
if self._clip is not None:
|
||||
self._clip = self._clip.to(self.device)
|
||||
if self._head is not None:
|
||||
self._head = self._head.to(self.device)
|
||||
return self
|
||||
|
||||
def compute(self, sample: dict) -> list[MetricResult]:
|
||||
...
|
||||
```
|
||||
|
||||
### Do not redirect other `~/.cache/...` dirs
|
||||
|
||||
If a third-party library hard-codes `~/.cache/<lib>/` and offers no
|
||||
override, document the exception in the metric's docstring. Forcing
|
||||
redirection by setting `os.environ` or patching `os.path.expanduser`
|
||||
is fragile and breaks user expectations of where the library's cache
|
||||
lives.
|
||||
|
||||
---
|
||||
|
||||
## 4) Recipe C: upstream-submodule-wrapped metric (vbench pattern)
|
||||
|
||||
Use this when the upstream benchmark ships Python code (`vbench/`,
|
||||
`MIND/`, etc.) that is not pip-installable cleanly. See
|
||||
`fastvideo/eval/metrics/vbench/__init__.py` for the worked example.
|
||||
|
||||
### 4.1 Pin the upstream as a submodule
|
||||
|
||||
```bash
|
||||
git submodule add <upstream-url> fastvideo/third_party/eval/<bench>
|
||||
cd fastvideo/third_party/eval/<bench>
|
||||
git checkout <pinned-sha>
|
||||
cd -
|
||||
git add .gitmodules fastvideo/third_party/eval/<bench>
|
||||
```
|
||||
|
||||
The submodule pulls under the standard `git submodule update --init
|
||||
--recursive` flow that users already run for kernel deps.
|
||||
|
||||
### 4.2 Bootstrap on `sys.path`
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/__init__.py
|
||||
from __future__ import annotations
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# fastvideo/eval/metrics/<bench>/__init__.py → ../../../../third_party/eval/<bench>
|
||||
_UPSTREAM = Path(__file__).resolve().parents[3] / "third_party" / "eval" / "<bench>"
|
||||
if _UPSTREAM.is_dir() and str(_UPSTREAM) not in sys.path:
|
||||
sys.path.insert(0, str(_UPSTREAM))
|
||||
```
|
||||
|
||||
We do not `pip install` the upstream because its egg-link/.pth would
|
||||
just re-do this `sys.path.insert`, and skipping the install also skips
|
||||
the upstream's `setup.py` (which often gates on a specific CUDA
|
||||
version).
|
||||
|
||||
### 4.3 Modern-dep compat: runtime shims
|
||||
|
||||
Upstream code pinned to e.g. `transformers==4.33.2`, `numpy<2`
|
||||
typically breaks against modern versions in 3-4 known places (API
|
||||
renames). Fix those at import time, in the same `__init__.py`:
|
||||
|
||||
```python
|
||||
def _install_compat_shims() -> None:
|
||||
# Example: transformers.modeling_utils API moved.
|
||||
try:
|
||||
import transformers.modeling_utils as _mu
|
||||
import transformers.pytorch_utils as _pu
|
||||
for _n in ("apply_chunking_to_forward",
|
||||
"find_pruneable_heads_and_indices",
|
||||
"prune_linear_layer"):
|
||||
if not hasattr(_mu, _n) and hasattr(_pu, _n):
|
||||
setattr(_mu, _n, getattr(_pu, _n))
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Example: numpy.lib.function_base.disp removed in numpy>=2.
|
||||
try:
|
||||
import types, numpy.lib as _nl
|
||||
if not hasattr(_nl, "function_base"):
|
||||
_stub = types.ModuleType("numpy.lib.function_base")
|
||||
_stub.disp = lambda *a, **k: None
|
||||
sys.modules["numpy.lib.function_base"] = _stub
|
||||
_nl.function_base = _stub
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
_install_compat_shims()
|
||||
```
|
||||
|
||||
For function-level patches that cannot be expressed as attribute
|
||||
writes (e.g. wrapping a model factory function), use a
|
||||
`sys.meta_path` finder that wraps the loader. See
|
||||
`_install_modeling_finetune_hook()` in
|
||||
`fastvideo/eval/metrics/vbench/__init__.py` for the pattern (about 30
|
||||
lines).
|
||||
|
||||
Why shims rather than `git apply` patches: patches go stale when the
|
||||
upstream SHA changes; shims are versioned Python code in our repo,
|
||||
they are grep-able, and they only run if the targeted module is
|
||||
imported.
|
||||
|
||||
### 4.4 Per-sub-metric files
|
||||
|
||||
Each sub-metric is a normal `BaseMetric` subclass that imports from
|
||||
the upstream:
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/<sub>/metric.py
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult
|
||||
|
||||
|
||||
@register("<bench>.<sub>")
|
||||
class YourSubMetric(BaseMetric):
|
||||
...
|
||||
|
||||
def setup(self) -> None:
|
||||
if self._model is not None:
|
||||
return
|
||||
# The sys.path bootstrap fired when fastvideo.eval.metrics.<bench>
|
||||
# was imported (which auto-discovery does before importing this
|
||||
# sub-package). Upstream imports just work:
|
||||
from <bench>.something import SomeModel
|
||||
...
|
||||
```
|
||||
|
||||
### 4.5 Conditional registration when the upstream is missing
|
||||
|
||||
If a user installed `fastvideo[eval]` but did not run `git submodule
|
||||
update --init`, `<bench>.*` metrics should not register. The
|
||||
auto-discovery walker imports each sub-package's `metric` module; have
|
||||
that import bail out cleanly:
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/__init__.py — at the bottom
|
||||
_AVAILABLE = (_UPSTREAM / "<bench>" / "__init__.py").is_file()
|
||||
```
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/<sub>/__init__.py
|
||||
from fastvideo.eval.metrics.<bench> import _AVAILABLE
|
||||
if _AVAILABLE:
|
||||
from .metric import YourSubMetric # noqa
|
||||
```
|
||||
|
||||
`fastvideo eval list` then reflects what the user actually has rather
|
||||
than what they could have.
|
||||
|
||||
### 4.6 Leave the upstream alone unless it blocks the metric
|
||||
|
||||
The upstream is pinned. If a metric works against the pinned SHA,
|
||||
leave the upstream files untouched. If it actively breaks against
|
||||
modern deps (the import-drift cases above), shim it. Avoid
|
||||
fastvideo-side forks of upstream code; they make patches go stale and
|
||||
parity drift.
|
||||
|
||||
---
|
||||
|
||||
## 5) Model checkpoints: `ensure_checkpoint`
|
||||
|
||||
Use `ensure_checkpoint(name, source, filename=None)` for any
|
||||
non-package weights. It resolves a local path, downloading on miss,
|
||||
with filelock safety across processes and SLURM ranks.
|
||||
|
||||
| `source` form | What happens |
|
||||
|---|---|
|
||||
| `"/abs/path/to/file.pth"` | passthrough, returned unchanged |
|
||||
| `"https://..."` | downloaded to `${FASTVIDEO_EVAL_CACHE}/models/<name>` via `huggingface_hub.http_get`, atomic rename, filelock |
|
||||
| `"org/repo"` (no `filename`) | `snapshot_download(repo_id)` → `~/.cache/huggingface/hub/` |
|
||||
| `"org/repo"` (with `filename`) | `hf_hub_download(repo_id, filename)` → `~/.cache/huggingface/hub/` |
|
||||
|
||||
`name` is only used as the local filename for URL sources. HF sources
|
||||
ignore it (HF manages its own cache key by content hash).
|
||||
|
||||
```python
|
||||
from fastvideo.eval.models import ensure_checkpoint
|
||||
|
||||
# URL: name matters
|
||||
ckpt = ensure_checkpoint(
|
||||
"amt-s.pth",
|
||||
source="https://huggingface.co/lalala125/AMT/resolve/main/amt-s.pth",
|
||||
)
|
||||
|
||||
# HF single file: name is decorative
|
||||
ckpt = ensure_checkpoint(
|
||||
"raft-things.pth", # ignored; HF cache uses repo+sha
|
||||
source="OpenGVLab/VBench_Used_Models",
|
||||
filename="raft-things.pth",
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 6) Declaring `dependencies`
|
||||
|
||||
Set `dependencies = ["pkg1", "pkg2"]` on your metric class with
|
||||
importable module names (not PyPI distribution names). The registry
|
||||
checks each via `importlib.util.find_spec` at instantiation time and
|
||||
raises a clean `ImportError` pointing the user at the right install
|
||||
extra:
|
||||
|
||||
```python
|
||||
class YourMetric(BaseMetric):
|
||||
dependencies = ["clip", "timm"] # importable as `import clip`, `import timm`
|
||||
```
|
||||
|
||||
If a dep is in `[project.optional-dependencies.eval-<group>]`, you do
|
||||
not need to do anything more. If it is a new dep, add it to that
|
||||
group in `pyproject.toml`.
|
||||
|
||||
---
|
||||
|
||||
## 7) Common gotchas
|
||||
|
||||
- The standard `git submodule update --init --recursive` is enough
|
||||
for a benchmark; do not write a `setup.sh`. Modern-dep compat goes
|
||||
into your `__init__.py` as runtime shims.
|
||||
- Do not modify upstream files on disk. The submodule should always
|
||||
match its pinned SHA. Compat lives in our `__init__.py`.
|
||||
- Do not pip-install the upstream. The egg-link is a glorified
|
||||
`sys.path.insert`, which we do directly in `__init__.py`.
|
||||
- Do not call `torch.hub.set_dir(...)` from your metric. It is done
|
||||
globally in `fastvideo/eval/__init__.py`.
|
||||
- Do not put cache-redirection env vars in your metric's `setup()`.
|
||||
By the time `setup()` runs, the library has likely already cached
|
||||
the default-location decision. Set env vars at package-init time.
|
||||
- Skip rather than raise when an input is missing. Use
|
||||
`self._skip(sample, reason)` for any expected-missing input. It
|
||||
returns a list of `MetricResult(score=None)` so other metrics in
|
||||
the same evaluator continue.
|
||||
- Watch for upstream re-registration conflicts. If the upstream uses
|
||||
a global registry (detectron2's `META_ARCH_REGISTRY`, MMCV, etc.),
|
||||
loading the same model twice in the same process will throw. The
|
||||
evaluator already loads each metric once; if you write a custom
|
||||
setup-then-call pattern, mirror that single-load discipline.
|
||||
|
||||
---
|
||||
|
||||
## 8) Training-time eval: keep evaluators hot, free caches between calls
|
||||
|
||||
When wiring eval into a training loop, the working pattern is:
|
||||
|
||||
1. Construct the `Evaluator` once and attach it to the pipeline
|
||||
(`self._eval = create_evaluator(...)`). Do not recreate it per
|
||||
validation round; that re-pays the model load cost.
|
||||
2. Save validation videos to disk (the diffusion path already does
|
||||
this). Pass paths to `evaluator.evaluate`, not in-memory tensors
|
||||
that share GPU memory with the training model.
|
||||
3. Run validation only on rank 0 of each sequence-parallel group.
|
||||
Gather paths from other ranks and let rank 0 score everything.
|
||||
4. After every `evaluate(...)` call, call
|
||||
`evaluator.release_cuda_memory()` in a `finally` block. That runs
|
||||
`gc.collect()` + `torch.cuda.empty_cache()` +
|
||||
`torch.cuda.ipc_collect()`. The eval model stays loaded; only
|
||||
transient activation buffers from the just-finished call get
|
||||
freed:
|
||||
|
||||
```python
|
||||
for video_path in batch:
|
||||
try:
|
||||
scores = self._eval.evaluate(video=load_video(video_path))
|
||||
finally:
|
||||
self._eval.release_cuda_memory()
|
||||
```
|
||||
|
||||
5. If memory pressure spikes (rare on H200), call
|
||||
`evaluator.unload()` to drop every metric reference and let the
|
||||
GPU memory be GC'd. `unload` is reversible:
|
||||
`evaluator.reload()` rebuilds the same metrics with the original
|
||||
config (re-paying the model load cost). Calling `evaluate`
|
||||
between `unload` and `reload` raises a clear `RuntimeError`.
|
||||
|
||||
For most metrics (sub-1 GB backbones, e.g. CLIP/DINO/RAFT/AMT) the
|
||||
eval model can stay co-resident with the training model in
|
||||
`transformer.eval()` mode without any swap. For larger ones
|
||||
(VideoScore2 at 14 GB), measure first; if it fits on the rank-0 GPU
|
||||
during validation (training model in eval mode means no
|
||||
grads/optimizer updates), keep it hot. If not, `unload` between
|
||||
rounds.
|
||||
|
||||
## 9) Local verification
|
||||
|
||||
Native and library-wrapped metrics: a single-GPU smoke is enough.
|
||||
|
||||
```python
|
||||
import torch
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
ev = create_evaluator(metrics=["<group>.<your_metric>"], device="cuda")
|
||||
video = torch.randn(1, 49, 3, 256, 256, device="cuda").clamp(0, 1)
|
||||
print(ev.evaluate(video=video))
|
||||
```
|
||||
|
||||
Submodule-wrapped metrics: also do a parity check against the
|
||||
upstream once. Clone upstream into a separate venv, run the same
|
||||
video through both, and expect an exact match on bit-deterministic
|
||||
metrics and ≤1% drift on backbone-heavy ones (driven by
|
||||
transformers/torch version differences).
|
||||
|
||||
For quick parity in CI: pin a tiny test video, record expected
|
||||
scores ± tolerance, and add a calibration test under
|
||||
`fastvideo/tests/eval/`.
|
||||
|
||||
---
|
||||
|
||||
## 10) When not to add a metric
|
||||
|
||||
- **Set-vs-set distribution metrics** (FVD, FID-style) do not fit
|
||||
`BaseMetric.compute(sample)` cleanly; they need a population.
|
||||
Adding them requires a stateful accumulator interface that does
|
||||
not exist yet. Open an issue first.
|
||||
- **Metrics requiring a single-GPU model larger than available
|
||||
memory.** Eval is not the place for tensor-parallel sharding;
|
||||
metrics are expected to fit on one GPU.
|
||||
- **Metrics that need `mmcv` with a conflicting CUDA ABI.** Document
|
||||
the affected sub-metrics as unsupported and skip them. Building
|
||||
isolation infrastructure (subprocess engine, per-metric venv) is
|
||||
out of scope.
|
||||
@@ -104,8 +104,9 @@ distillation, self-forcing, VSA, VMoBA, performance benchmarks, and API server t
|
||||
8. If all Full Suite tests pass and all merge conditions are met (approval, valid title,
|
||||
pre-commit green, fastcheck green, no draft, no conflicts), Mergify squash-merges to
|
||||
`main` automatically. Your branch is deleted.
|
||||
9. If a Full Suite test fails, Mergify removes the `ready` label and posts a comment with a
|
||||
link to the Buildkite build. Fix the issue, push, and comment `/merge` again.
|
||||
9. If a Full Suite test fails, check the Buildkite build log for the failing step. Fix the
|
||||
issue, push, and comment `/merge` again. You can also re-run individual failed tests
|
||||
with `/test <name>` — see below.
|
||||
|
||||
!!! note
|
||||
Only contributors with write permission to the repository can trigger slash commands.
|
||||
@@ -149,10 +150,15 @@ Comment on your PR to trigger specific tests independently of the auto-merge flo
|
||||
/test vmoba # VMoBA inference tests
|
||||
/test performance # Performance benchmarks
|
||||
/test api # API server integration tests
|
||||
/test pre-commit # Pre-commit checks on PR code
|
||||
```
|
||||
|
||||
The workflow reacts with a 🚀 emoji to confirm the command was received.
|
||||
|
||||
When you re-run an individual test with `/test <name>`, the new result overwrites the
|
||||
original failed check (same Buildkite check name). Once all tests in a tier pass, the
|
||||
`fastcheck-passed` or `full-suite-passed` status is automatically updated.
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
@@ -199,9 +205,8 @@ Mergify removes the `needs-rebase` label automatically once conflicts are resolv
|
||||
|
||||
### Full Suite failed after `/merge`
|
||||
|
||||
The Full Suite found a regression. Mergify removes the `ready` label and posts a comment
|
||||
linking to the Buildkite build. Check the failing step's output for assertion errors or
|
||||
tracebacks.
|
||||
The Full Suite found a regression. Check the failing Buildkite step's output for assertion
|
||||
errors or tracebacks.
|
||||
|
||||
Common causes:
|
||||
|
||||
|
||||
@@ -0,0 +1,449 @@
|
||||
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."
|
||||
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."
|
||||
|
||||
surfaces:
|
||||
fastvideo_args:
|
||||
moved:
|
||||
model_path: generator.model_path
|
||||
workload_type: generator.pipeline.workload_type
|
||||
distributed_executor_backend: generator.engine.execution_backend
|
||||
trust_remote_code: generator.trust_remote_code
|
||||
revision: generator.revision
|
||||
num_gpus: generator.engine.num_gpus
|
||||
tp_size: generator.engine.parallelism.tp_size
|
||||
sp_size: generator.engine.parallelism.sp_size
|
||||
hsdp_replicate_dim: generator.engine.parallelism.hsdp_replicate_dim
|
||||
hsdp_shard_dim: generator.engine.parallelism.hsdp_shard_dim
|
||||
dist_timeout: generator.engine.parallelism.dist_timeout
|
||||
lora_path: generator.pipeline.components.lora_path
|
||||
dit_cpu_offload: generator.engine.offload.dit
|
||||
use_fsdp_inference: generator.engine.use_fsdp_inference
|
||||
dit_layerwise_offload: generator.engine.offload.dit_layerwise
|
||||
text_encoder_cpu_offload: generator.engine.offload.text_encoder
|
||||
image_encoder_cpu_offload: generator.engine.offload.image_encoder
|
||||
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.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
|
||||
override_text_encoder_safetensors: generator.pipeline.components.text_encoder_weights
|
||||
override_text_encoder_quant: generator.engine.quantization.text_encoder_quant
|
||||
override_transformer_cls_name: generator.pipeline.components.override_transformer_cls_name
|
||||
init_weights_from_safetensors: generator.pipeline.components.transformer_weights
|
||||
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
|
||||
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."
|
||||
inference_mode: "Legacy boolean mirror of mode; kept only through adapters while FastVideoArgs remains."
|
||||
lora_nickname: "Legacy adapter-selection surface pending LoRA API cleanup."
|
||||
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
|
||||
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
|
||||
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
|
||||
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
|
||||
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
|
||||
private_only:
|
||||
ray_placement_group: "Ray deployment-only field."
|
||||
ray_runtime_env: "Ray deployment-only field."
|
||||
internal_only:
|
||||
pipeline_config: "Legacy internal carrier object."
|
||||
preprocess_config: "Legacy preprocess carrier object."
|
||||
moba_config: "Derived runtime config loaded from moba_config_path."
|
||||
model_paths: "Runtime bookkeeping."
|
||||
model_loaded: "Runtime bookkeeping."
|
||||
|
||||
pipeline_config_base:
|
||||
moved:
|
||||
pipeline_config_path: generator.pipeline.components.pipeline_config_path
|
||||
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."
|
||||
dit_precision: "Precision override pending dedicated typed component precision design."
|
||||
upsampler_precision: "Precision override pending dedicated typed component precision design."
|
||||
vae_precision: "Precision override pending dedicated typed component precision design."
|
||||
image_encoder_precision: "Precision override pending dedicated typed component precision design."
|
||||
text_encoder_precisions: "Precision override pending dedicated typed component precision design."
|
||||
internal_only:
|
||||
dit_config: "Legacy internal component config object."
|
||||
upsampler_config: "Legacy internal component config object."
|
||||
vae_config: "Legacy internal component config object."
|
||||
image_encoder_config: "Legacy internal component config object."
|
||||
text_encoder_configs: "Legacy internal component config object."
|
||||
preprocess_text_funcs: "Internal text preprocessing hooks."
|
||||
postprocess_text_funcs: "Internal text postprocessing hooks."
|
||||
|
||||
pipeline_config_extensions:
|
||||
preset_owned:
|
||||
conditioning_strategy:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
max_num_conditional_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
min_num_conditional_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
sigma_conditional:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
sigma_data:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
state_ch:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
state_t:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
text_encoder_class:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
autoregressive_chunk_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
autoregressive_overlap_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
cfg_behavior:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_camera_rotation:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_movement_distance:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_negative_prompt:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
default_trajectory_type:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
filter_points_threshold:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
fps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
frame_buffer_max:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
moge_model_name:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
noise_aug_strength:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
num_frames:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
offload_moge_after_depth:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
use_moge_depth:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
video_resolution:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
text_encoder_crop_start:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V480PStepDistilledConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V720PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15SR1080PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V480PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V720PConfig
|
||||
- fastvideo.configs.pipelines.hyworld.HYWorldConfig
|
||||
- fastvideo.configs.pipelines.hyworld.Hunyuan15T2V480PConfig
|
||||
text_encoder_max_lengths:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V480PStepDistilledConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15I2V720PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15SR1080PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V480PConfig
|
||||
- fastvideo.configs.pipelines.hunyuan15.Hunyuan15T2V720PConfig
|
||||
- fastvideo.configs.pipelines.hyworld.HYWorldConfig
|
||||
- fastvideo.configs.pipelines.hyworld.Hunyuan15T2V480PConfig
|
||||
precision:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.lingbotworld.LingBotWorldI2V480PConfig
|
||||
- fastvideo.configs.pipelines.lingbotworld.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_1_3B_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_1_T2V_480P_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameBaseI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWan2_2_T2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WANV2VConfig
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_T2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.WanI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanI2V720PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V720PConfig
|
||||
warp_denoising_step:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.lingbotworld.LingBotWorldI2V480PConfig
|
||||
- fastvideo.configs.pipelines.lingbotworld.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionI2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2VConfig
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_14B_Config
|
||||
- fastvideo.configs.pipelines.turbodiffusion.TurboDiffusionT2V_1_3B_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_1_T2V_480P_Config
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameBaseI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWan2_2_T2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.SelfForcingWanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WANV2VConfig
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_I2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_T2V_A14B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.WanI2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanI2V720PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V480PConfig
|
||||
- fastvideo.configs.pipelines.wan.WanT2V720PConfig
|
||||
bsa_cdf_threshold:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_chunk_k:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_chunk_q:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_params:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
bsa_sparsity:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
enable_bsa:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
enable_kv_cache:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
enhance_hf:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
offload_kv_cache:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
t_thresh:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
use_distill:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
scheduler_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
text_encoder_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
tokenizer_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
transformer_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
vae_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
expand_timesteps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
- fastvideo.configs.pipelines.wan.Wan2_2_TI2V_5B_Config
|
||||
context_noise:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
num_frames_per_block:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
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 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."
|
||||
vocoder_config: "Legacy internal component config object."
|
||||
vocoder_precision: "Precision override pending dedicated component precision design."
|
||||
|
||||
sampling_param_base:
|
||||
moved:
|
||||
image_path: request.inputs.image_path
|
||||
pil_image: request.inputs.pil_image
|
||||
video_path: request.inputs.video_path
|
||||
mouse_cond: request.inputs.mouse_cond
|
||||
keyboard_cond: request.inputs.keyboard_cond
|
||||
grid_sizes: request.inputs.grid_sizes
|
||||
pose: request.inputs.pose
|
||||
c2ws_plucker_emb: request.inputs.c2ws_plucker_emb
|
||||
refine_from: request.inputs.refine_from
|
||||
stage1_video: request.inputs.stage1_video
|
||||
prompt: request.prompt
|
||||
negative_prompt: request.negative_prompt
|
||||
prompt_path: request.inputs.prompt_path
|
||||
output_path: request.output.output_path
|
||||
output_video_name: request.output.output_video_name
|
||||
num_videos_per_prompt: request.sampling.num_videos_per_prompt
|
||||
seed: request.sampling.seed
|
||||
num_frames: request.sampling.num_frames
|
||||
height: request.sampling.height
|
||||
width: request.sampling.width
|
||||
height_sr: request.sampling.height_sr
|
||||
width_sr: request.sampling.width_sr
|
||||
fps: request.sampling.fps
|
||||
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
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
save_video: request.output.save_video
|
||||
return_frames: request.output.return_frames
|
||||
return_trajectory_latents: request.runtime.return_trajectory_latents
|
||||
return_trajectory_decoded: request.runtime.return_trajectory_decoded
|
||||
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
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
|
||||
sampling_param_extensions: {}
|
||||
|
||||
openai_image_request:
|
||||
kept:
|
||||
model: "HTTP adapter model-routing field."
|
||||
response_format: "HTTP adapter response formatting field."
|
||||
output_format: "HTTP adapter output-format field."
|
||||
background: "HTTP adapter output-format field."
|
||||
quality: "Compatibility field currently accepted by the adapter."
|
||||
style: "Compatibility field currently accepted by the adapter."
|
||||
user: "Compatibility field currently accepted by the adapter."
|
||||
moved:
|
||||
prompt: request.prompt
|
||||
n: request.sampling.num_videos_per_prompt
|
||||
size:
|
||||
target: request.sampling.width,height
|
||||
note: "Adapter parses OpenAI size strings as WIDTHxHEIGHT and forwards width then height."
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
seed: request.sampling.seed
|
||||
negative_prompt: request.negative_prompt
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
|
||||
openai_video_request:
|
||||
kept:
|
||||
model: "HTTP adapter model-routing field."
|
||||
moved:
|
||||
prompt: request.prompt
|
||||
input_reference: request.inputs.image_path
|
||||
reference_url: request.inputs.image_path
|
||||
size:
|
||||
target: request.sampling.width,height
|
||||
note: "Adapter parses OpenAI size strings as WIDTHxHEIGHT and forwards width then height."
|
||||
fps: request.sampling.fps
|
||||
num_frames: request.sampling.num_frames
|
||||
seed: request.sampling.seed
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
negative_prompt: request.negative_prompt
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
output_path: request.output.output_path
|
||||
compatibility_only:
|
||||
seconds:
|
||||
target: request.sampling.num_frames
|
||||
note: "HTTP adapter duration convenience field. If num_frames is omitted, the adapter computes num_frames = fps * seconds."
|
||||
|
||||
cli:
|
||||
notes:
|
||||
- "CLI parity is checked against the actual generate/serve parser dest sets."
|
||||
- "The inventory tracks parser dest names, excluding argparse's implicit help action."
|
||||
- "The refactored inference CLI is config-only: subcommands expose only --config, and any additional CLI input must use dotted override paths."
|
||||
generate:
|
||||
explicit_local_fields:
|
||||
- config
|
||||
expected_dests:
|
||||
- config
|
||||
serve:
|
||||
explicit_local_fields:
|
||||
- config
|
||||
expected_dests:
|
||||
- config
|
||||
@@ -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`.
|
||||
@@ -16,7 +16,8 @@ Both models are trained on **61×448×832** resolution but support generating vi
|
||||
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN \
|
||||
fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml
|
||||
```
|
||||
|
||||
## 🗂️ Dataset
|
||||
@@ -85,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.
|
||||
|
||||
@@ -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
|
||||
@@ -455,6 +450,6 @@ User: generator.generate_video(prompt, ...)
|
||||
`fastvideo/pipelines/stages/`, implement `forward()`, optionally
|
||||
implement `verify_input()`/`verify_output()`.
|
||||
|
||||
7. **Verify** — Run `fastvideo generate --model-path <path> --prompt
|
||||
"test" --num-inference-steps 2` to confirm the pipeline loads and
|
||||
generates output.
|
||||
7. **Verify** — Run `fastvideo generate --config <config.yaml>` with a
|
||||
minimal nested config to confirm the pipeline loads and generates
|
||||
output.
|
||||
|
||||
+42
-81
@@ -1,71 +1,29 @@
|
||||
# FastVideo CLI Inference
|
||||
|
||||
The FastVideo CLI exposes the same core inference controls as the Python API.
|
||||
The FastVideo CLI is config-first. Inference runs are driven by a nested JSON or
|
||||
YAML config, with optional dotted-path overrides on the command line. The
|
||||
contract matches training: use an explicit subcommand plus `--config`, then add
|
||||
any dotted overrides you need.
|
||||
|
||||
## Basic Usage
|
||||
|
||||
Use either:
|
||||
|
||||
1. `--model-path` + `--prompt`
|
||||
2. `--model-path` + `--prompt-txt` (batch prompts, one line per prompt)
|
||||
3. `--config` (JSON/YAML)
|
||||
|
||||
```bash
|
||||
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--prompt "A cat playing with a ball of yarn"
|
||||
fastvideo generate --config config.yaml
|
||||
fastvideo serve --config serve.yaml
|
||||
```
|
||||
|
||||
```bash
|
||||
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--prompt-txt prompts.txt
|
||||
```
|
||||
|
||||
You cannot provide both `--prompt` and `--prompt-txt` in the same run.
|
||||
|
||||
## View All Arguments
|
||||
|
||||
```bash
|
||||
fastvideo generate --help
|
||||
```
|
||||
|
||||
Arguments come from:
|
||||
The subcommands intentionally expose only `--config`. Any per-run CLI changes
|
||||
must use dotted override paths such as:
|
||||
|
||||
- FastVideo runtime args (`FastVideoArgs`)
|
||||
- Sampling args (`SamplingParam`)
|
||||
- Pipeline config args (`PipelineConfig`)
|
||||
|
||||
## Common Arguments
|
||||
|
||||
### Parallelism
|
||||
|
||||
- `--num-gpus`
|
||||
- `--sp-size`
|
||||
- `--tp-size`
|
||||
|
||||
### Sampling
|
||||
|
||||
- `--num-frames`
|
||||
- `--height` / `--width`
|
||||
- `--num-inference-steps`
|
||||
- `--guidance-scale`
|
||||
- `--seed`
|
||||
- `--negative-prompt`
|
||||
|
||||
### Output
|
||||
|
||||
- `--output-path`
|
||||
- `--save-video` / `--no-save-video`
|
||||
- `--return-frames`
|
||||
|
||||
### Offloading and Performance
|
||||
|
||||
- `--dit-layerwise-offload`
|
||||
- `--use-fsdp-inference`
|
||||
- `--text-encoder-cpu-offload`
|
||||
- `--image-encoder-cpu-offload`
|
||||
- `--vae-cpu-offload`
|
||||
- `--enable-torch-compile`
|
||||
- `--torch-compile-kwargs`
|
||||
- `--generator.engine.num_gpus 2`
|
||||
- `--request.sampling.seed 42`
|
||||
- `--server.port 9000`
|
||||
|
||||
## Using Config Files
|
||||
|
||||
@@ -73,50 +31,53 @@ Arguments come from:
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Config files can be JSON or YAML. CLI flags override config-file values.
|
||||
Config files can be JSON or YAML. Dotted CLI overrides take precedence over
|
||||
config-file values.
|
||||
|
||||
Example `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model_path: "FastVideo/FastHunyuan-diffusers"
|
||||
prompt: "A capybara lounging in a hammock"
|
||||
output_path: "outputs/"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
tp_size: 1
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
dit_precision: "bf16"
|
||||
vae_precision: "fp16"
|
||||
vae_tiling: true
|
||||
vae_sp: true
|
||||
enable_torch_compile: false
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
engine:
|
||||
num_gpus: 2
|
||||
parallelism:
|
||||
sp_size: 2
|
||||
tp_size: 1
|
||||
request:
|
||||
prompt: A capybara lounging in a hammock
|
||||
sampling:
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
output:
|
||||
output_path: outputs/
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- Use `dit_precision` / `vae_precision` (not `precision`).
|
||||
- Nested config objects are supported, for example `vae_config` and
|
||||
`dit_config`.
|
||||
- `generator` and `request` are the top-level keys for generation configs.
|
||||
- `serve` configs use `generator`, `server`, and optional `default_request`.
|
||||
- Prompt text files belong under `request.inputs.prompt_path`.
|
||||
|
||||
## Examples
|
||||
|
||||
Simple generation:
|
||||
|
||||
```bash
|
||||
fastvideo generate \
|
||||
--model-path FastVideo/FastHunyuan-diffusers \
|
||||
--prompt "A cat playing with a ball of yarn" \
|
||||
--num-frames 45 --height 720 --width 1280 \
|
||||
--num-inference-steps 6 --seed 1024 \
|
||||
--output-path outputs/
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Config + CLI override:
|
||||
Config + dotted override:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml --prompt "A panda skiing at sunset"
|
||||
fastvideo generate --config config.yaml --request.prompt "A panda skiing at sunset"
|
||||
```
|
||||
|
||||
Helper wrapper with positional config path:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/run.sh scripts/inference/inference_wan.yaml
|
||||
```
|
||||
|
||||
@@ -73,32 +73,40 @@ if __name__ == '__main__':
|
||||
|
||||
## JSON/YAML Config Files (CLI)
|
||||
|
||||
The CLI supports `--config` with JSON or YAML. Command-line arguments override
|
||||
config file values.
|
||||
By default, `fastvideo generate` uses `return_frames=false` unless you set
|
||||
`--return-frames` (or `return_frames: true` in config).
|
||||
The inference CLI is config-first. Use an explicit subcommand with `--config`,
|
||||
then apply optional dotted overrides on top, matching the training CLI style.
|
||||
By default, CLI generation uses `return_frames=false` unless you set
|
||||
`request.output.return_frames: true` in config or via a dotted override.
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Use CLI argument names as keys (underscore or hyphen is accepted). Example:
|
||||
Example nested config:
|
||||
|
||||
```yaml
|
||||
model_path: "FastVideo/FastHunyuan-diffusers"
|
||||
prompt: "A capybara relaxing in a hammock"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
dit_precision: "bf16"
|
||||
vae_precision: "fp16"
|
||||
vae_tiling: true
|
||||
vae_sp: true
|
||||
enable_torch_compile: false
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
engine:
|
||||
num_gpus: 2
|
||||
parallelism:
|
||||
sp_size: 2
|
||||
request:
|
||||
prompt: A capybara relaxing in a hammock
|
||||
sampling:
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
output:
|
||||
output_path: outputs/
|
||||
```
|
||||
|
||||
Override individual values from the CLI with dotted paths:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml --request.sampling.seed 42
|
||||
```
|
||||
|
||||
## Performance Optimization
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
# GEN3C: 3D-Informed Camera-Controlled Video Generation
|
||||
|
||||
[GEN3C](https://arxiv.org/abs/2503.03751) is NVIDIA's Cosmos-7B-based video model for camera-controlled generation from a single image. The FastVideo integration supports the GEN3C I2V workflow, including 3D cache conditioning and tokenizer-based conditioning latents.
|
||||
|
||||
## Key Features
|
||||
|
||||
- **Camera trajectory control**: `left/right/up/down/zoom_in/zoom_out/clockwise/counterclockwise`
|
||||
- **3D cache conditioning**: depth prediction -> point cloud cache -> forward warping -> latent conditioning
|
||||
- **Single-image to video generation**: 121-frame generation with camera motion
|
||||
- **Official raw checkpoint conversion**: `model.pt` -> Diffusers/FastVideo layout
|
||||
|
||||
## Model Sources
|
||||
|
||||
- Official raw checkpoint (not Diffusers): `nvidia/GEN3C-Cosmos-7B`
|
||||
- Diffusers-format checkpoint: `FastVideo/GEN3C-Cosmos-7B-Diffusers`
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Install MoGe:
|
||||
|
||||
```bash
|
||||
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:
|
||||
|
||||
```bash
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Option A: Use Diffusers-format weights directly
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/basic_gen3c.py \
|
||||
--model_path FastVideo/GEN3C-Cosmos-7B-Diffusers \
|
||||
--image_path /path/to/input.png \
|
||||
--prompt "" \
|
||||
--trajectory left \
|
||||
--movement_distance 0.3 \
|
||||
--camera_rotation center_facing \
|
||||
--num_inference_steps 35 \
|
||||
--guidance_scale 1.0 \
|
||||
--output_path outputs_video/gen3c_output.mp4
|
||||
```
|
||||
|
||||
### Option B: Convert official raw checkpoint locally
|
||||
|
||||
1. Download:
|
||||
|
||||
```bash
|
||||
huggingface-cli download nvidia/GEN3C-Cosmos-7B --local-dir official_weights/GEN3C-Cosmos-7B
|
||||
```
|
||||
|
||||
1. Convert:
|
||||
|
||||
```bash
|
||||
python scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py \
|
||||
--source official_weights/GEN3C-Cosmos-7B/model.pt \
|
||||
--output converted_weights/GEN3C-Cosmos-7B
|
||||
```
|
||||
|
||||
1. Run:
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/basic_gen3c.py \
|
||||
--model_path converted_weights/GEN3C-Cosmos-7B \
|
||||
--image_path /path/to/input.png \
|
||||
--prompt "" \
|
||||
--trajectory left \
|
||||
--movement_distance 0.3 \
|
||||
--camera_rotation center_facing \
|
||||
--num_inference_steps 35 \
|
||||
--guidance_scale 1.0 \
|
||||
--output_path outputs_video/gen3c_output.mp4
|
||||
```
|
||||
|
||||
## FastVideo Defaults
|
||||
|
||||
GEN3C defaults in FastVideo:
|
||||
|
||||
- `height=704`, `width=1280`
|
||||
- `num_frames=121`
|
||||
- `num_inference_steps=35`
|
||||
- `guidance_scale=1.0`
|
||||
- `fps=24`
|
||||
|
||||
These values are defined in:
|
||||
|
||||
- `fastvideo/pipelines/basic/gen3c/profiles.py`
|
||||
- `fastvideo/configs/pipelines/gen3c.py`
|
||||
|
||||
and align with the official GEN3C inference defaults in:
|
||||
|
||||
- `tmp/GEN3C/cosmos_predict1/diffusion/inference/inference_utils.py`
|
||||
|
||||
## Scheduler Note
|
||||
|
||||
The converted GEN3C Diffusers layout may include a FlowMatch scheduler config, but GEN3C denoising uses EDM preconditioning behavior. FastVideo's GEN3C pipeline enforces an EDM scheduler at runtime for parity with official inference behavior.
|
||||
|
||||
Implementation path:
|
||||
|
||||
- `fastvideo/pipelines/basic/gen3c/gen3c_pipeline.py`
|
||||
|
||||
## 3D Cache Conditioning Path
|
||||
|
||||
FastVideo GEN3C conditioning stage performs:
|
||||
|
||||
1. MoGe depth estimation from input image
|
||||
2. 3D cache initialization
|
||||
3. Camera trajectory generation
|
||||
4. Forward rendering of warped frames + masks
|
||||
5. VAE/tokenizer encoding of conditioning buffers
|
||||
6. Denoising with condition mask + condition pose channels
|
||||
|
||||
Main implementation:
|
||||
|
||||
- `fastvideo/pipelines/basic/gen3c/gen3c_pipeline.py`
|
||||
- `fastvideo/pipelines/basic/gen3c/cache_3d.py`
|
||||
- `fastvideo/pipelines/basic/gen3c/depth_estimation.py`
|
||||
- `fastvideo/models/vaes/gen3c_tokenizer_vae.py`
|
||||
|
||||
## References
|
||||
|
||||
- [GEN3C Paper](https://arxiv.org/abs/2503.03751)
|
||||
- [Official Repository](https://github.com/nv-tlabs/GEN3C)
|
||||
- [Official Checkpoint (raw)](https://huggingface.co/nvidia/GEN3C-Cosmos-7B)
|
||||
@@ -73,6 +73,7 @@ pipeline initialization and sampling.
|
||||
| Matrix Game 2.0 Base | `FastVideo/Matrix-Game-2.0-Base-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 GTA | `FastVideo/Matrix-Game-2.0-GTA-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Matrix Game 2.0 TempleRun | `FastVideo/Matrix-Game-2.0-TempleRun-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
@@ -85,6 +86,11 @@ The authoritative source for model-ID recognition is
|
||||
`fastvideo/registry.py`. If a model ID is registered there, FastVideo can
|
||||
resolve default pipeline and sampling configuration for it.
|
||||
|
||||
**Note (GEN3C)**: The official `nvidia/GEN3C-Cosmos-7B` repo provides a raw
|
||||
`model.pt` checkpoint. Use a Diffusers-format repo (for example,
|
||||
`FastVideo/GEN3C-Cosmos-7B-Diffusers`) or convert locally with
|
||||
`scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py`.
|
||||
|
||||
## Special requirements
|
||||
|
||||
### Sliding Tile Attention
|
||||
|
||||
@@ -28,6 +28,11 @@ For an example running DMD+VSA inference:
|
||||
python examples/inference/basic/basic_dmd.py
|
||||
```
|
||||
|
||||
For the typed config/request path added during the inference API refactor:
|
||||
```
|
||||
python examples/inference/basic/basic_dmd_new_api.py
|
||||
```
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
PipelineSelection,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2_typed"
|
||||
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
model_name = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
# PR 2 still routes a few advanced inference knobs through the
|
||||
# compatibility bridge until they get first-class typed fields.
|
||||
pipeline=PipelineSelection(
|
||||
experimental={
|
||||
"VSA_sparsity": 0.8,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
|
||||
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
|
||||
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
|
||||
"LED umbrella. Steam rises from a street food cart, and a cat darts "
|
||||
"across the screen. Raindrops are visible on the camera lens, creating "
|
||||
"a cinematic bokeh effect."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
end_time = time.perf_counter()
|
||||
gen_time = end_time - start_time
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently "
|
||||
"in the breeze, enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low angle, steady "
|
||||
"tracking shot, cinematic."
|
||||
)
|
||||
request2 = GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result2 = generator.generate(request2)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"First output written to: {result.video_path}")
|
||||
print(f"Time taken to generate video2: {gen_time2} seconds")
|
||||
print(f"Second output written to: {result2.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,109 @@
|
||||
"""
|
||||
GEN3C: 3D-aware camera-controlled video generation.
|
||||
|
||||
This example generates a video from a single input image with camera control.
|
||||
The pipeline uses MoGe depth estimation, 3D point cloud forward warping,
|
||||
and the GEN3C diffusion model.
|
||||
|
||||
Requirements:
|
||||
1. Install MoGe:
|
||||
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:
|
||||
huggingface-cli download nvidia/GEN3C-Cosmos-7B --local-dir official_weights/GEN3C-Cosmos-7B
|
||||
python scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py \
|
||||
--source ./official_weights/GEN3C-Cosmos-7B/model.pt \
|
||||
--output ./converted_weights/GEN3C-Cosmos-7B \
|
||||
--components-source nvidia/Cosmos-Predict2-2B-Video2World
|
||||
3. Provide an input image for 3D-conditioned generation.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="GEN3C video generation")
|
||||
parser.add_argument("--model_path",
|
||||
type=str,
|
||||
default="converted_weights/GEN3C-Cosmos-7B")
|
||||
parser.add_argument("--image_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Input image for 3D cache conditioning")
|
||||
parser.add_argument("--prompt",
|
||||
type=str,
|
||||
default="A slow camera pan over a sunlit landscape.")
|
||||
parser.add_argument(
|
||||
"--negative_prompt",
|
||||
type=str,
|
||||
default=(
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--trajectory",
|
||||
type=str,
|
||||
default="left",
|
||||
choices=[
|
||||
"left", "right", "up", "down", "zoom_in",
|
||||
"zoom_out", "clockwise", "counterclockwise", "none"
|
||||
])
|
||||
parser.add_argument("--movement_distance", type=float, default=0.3)
|
||||
parser.add_argument("--camera_rotation",
|
||||
type=str,
|
||||
default="center_facing",
|
||||
choices=[
|
||||
"center_facing", "no_rotation",
|
||||
"trajectory_aligned"
|
||||
])
|
||||
parser.add_argument("--height", type=int, default=704)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=121)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=35)
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0)
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs_video/gen3c.mp4")
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
args = parser.parse_args()
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
video = generator.generate_video(
|
||||
args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
image_path=args.image_path,
|
||||
trajectory_type=args.trajectory,
|
||||
movement_distance=args.movement_distance,
|
||||
camera_rotation=args.camera_rotation,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
fps=24,
|
||||
seed=args.seed,
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
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"
|
||||
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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
"""Generate one LTX2 video and score it with VBench metrics.
|
||||
|
||||
The generation block is the same as
|
||||
``examples/inference/basic/basic_ltx2.py`` — same prompt, same model,
|
||||
same shape, same num_frames. After ``shutdown()`` the script loads the
|
||||
mp4 back, builds a single :class:`fastvideo.eval.Evaluator`, and runs
|
||||
the prompt-aware VBench subset that's meaningful for an arbitrary
|
||||
text→video sample.
|
||||
|
||||
The first run downloads CLIP / DINO / RAFT / AMT / ViCLIP / MUSIQ
|
||||
weights to ``~/.cache/fastvideo/eval/`` (~few GB total).
|
||||
|
||||
GPU memory caveat
|
||||
-----------------
|
||||
Scoring 1088×1920×121 with all 8 metrics needs a dedicated GPU (~80 GB).
|
||||
On a shared GPU, ``vbench.motion_smoothness`` (AMT correlation volume)
|
||||
will OOM — its memory autoscale reads ``total_memory`` rather than
|
||||
``mem_get_info()`` free memory and therefore underestimates the
|
||||
required scale-down. Drop ``motion_smoothness`` from ``METRICS`` if
|
||||
sharing, or run on a smaller-resolution generation.
|
||||
"""
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import Evaluator
|
||||
from fastvideo.eval.io import build_eval_kwargs
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
# VBench sub-metrics meaningful for an arbitrary text→video sample
|
||||
# (just the generated frames, optionally fps + the source prompt).
|
||||
# Structured-prompt metrics (vbench.color, vbench.multiple_objects,
|
||||
# vbench.scene, ...) are excluded — they need prompts built to a
|
||||
# specific schema.
|
||||
METRICS = [
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.background_consistency", # DINO on background patches
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
]
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# ----- generation (matches examples/inference/basic/basic_ltx2.py) -----
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
)
|
||||
generator.shutdown()
|
||||
# Free residual CUDA memory the generator left behind so the
|
||||
# evaluator can grab the largest possible workspace for AMT/RAFT.
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# ----- scoring -----
|
||||
print(f"\n[eval] building evaluator: {METRICS}")
|
||||
evaluator = Evaluator(metrics=METRICS)
|
||||
|
||||
# LTX2 outputs at 24 fps by default.
|
||||
sample = build_eval_kwargs({"prompt": PROMPT}, output_path, fps=24.0)
|
||||
print(f"[eval] running ({sample['video'].shape[1]} frames @ 24 fps)...")
|
||||
results = evaluator.evaluate(**sample)
|
||||
|
||||
print("\n=== VBench scores ===")
|
||||
for name in METRICS:
|
||||
r = results[name]
|
||||
if r.score is None:
|
||||
reason = r.details.get("skipped", "no score")
|
||||
print(f" {name}: SKIPPED ({reason})")
|
||||
else:
|
||||
print(f" {name}: {r.score:.4f}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,157 @@
|
||||
"""End-to-end Physics-IQ: dataset → generate → score → aggregate.
|
||||
|
||||
Generates one video per take-1 scenario with LTX2 (using the scenario
|
||||
caption as the prompt), scores each generated video against the take-1
|
||||
reference and the take-2 "physical-variance" reference, and prints
|
||||
aggregate scores using :meth:`PhysicsIQMetric.aggregate_components` —
|
||||
the official scoring recipe from the upstream benchmark.
|
||||
|
||||
Reference videos / masks / switch-frames auto-fetch on first miss into
|
||||
``${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/``; pass ``--dataset-root``
|
||||
to point at a pre-downloaded mirror instead.
|
||||
|
||||
Quick smoke run on 4 scenarios across 2 GPUs::
|
||||
|
||||
python examples/inference/eval/bench_physics_iq.py \\
|
||||
--limit 4 --num-gpus 2 \\
|
||||
--videos-dir outputs_video/physics_iq_smoke
|
||||
|
||||
Re-score existing generations without regenerating::
|
||||
|
||||
python examples/inference/eval/bench_physics_iq.py \\
|
||||
--videos-dir outputs_video/physics_iq_smoke \\
|
||||
--skip-generation
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator, get_metric
|
||||
from fastvideo.eval.datasets import get_dataset
|
||||
|
||||
|
||||
def _expected_filename(row: dict) -> str:
|
||||
"""Filename Physics-IQ expects for the generated video for *row*.
|
||||
|
||||
Uses the dataset's own ``expected_gen_filename`` annotation so the
|
||||
output filenames match the benchmark's manifest convention.
|
||||
"""
|
||||
return row["auxiliary_info"]["expected_gen_filename"]
|
||||
|
||||
|
||||
def _generate_videos(rows: list[dict], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
todo = [(row, videos_dir / _expected_filename(row)) for row in rows]
|
||||
todo = [(row, out) for (row, out) in todo if not out.is_file()]
|
||||
if not todo:
|
||||
print(f"[gen] all {len(rows)} videos already present; skipping.")
|
||||
return
|
||||
|
||||
print(f"[gen] {len(todo)}/{len(rows)} scenarios to render with {model} "
|
||||
f"({num_frames}x{height}x{width})...")
|
||||
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
|
||||
try:
|
||||
for row, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=row["prompt"], output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
)
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dataset-root", type=Path, default=None,
|
||||
help="Path to a pre-downloaded Physics-IQ release. "
|
||||
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
|
||||
"auto-fetching missing assets from the public bucket.")
|
||||
p.add_argument("--videos-dir", type=Path,
|
||||
default=Path("outputs_video/bench_physics_iq"),
|
||||
help="Where to read/write generated videos.")
|
||||
p.add_argument("--limit", type=int, default=None,
|
||||
help="Truncate to first N scenarios for smoke runs.")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the text→video generator to use.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Re-score existing videos under --videos-dir.")
|
||||
p.add_argument("--scores-out", type=Path, default=None,
|
||||
help="Where to write per-scenario scores (JSON). "
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
args = p.parse_args()
|
||||
|
||||
# 1. Walk the Physics-IQ corpus. Pass --limit to the dataset
|
||||
# constructor so auto-download only fetches the assets we'll use.
|
||||
ds = get_dataset("physics_iq", dataset_root=args.dataset_root, limit=args.limit)
|
||||
rows = list(ds)
|
||||
print(f"[load] Physics-IQ: {len(rows)} scenarios from {ds.dataset_dir}")
|
||||
|
||||
# 2. Generate (or reuse) one mp4 per scenario.
|
||||
if not args.skip_generation:
|
||||
_generate_videos(
|
||||
rows, args.videos_dir, args.model, args.num_gpus,
|
||||
args.num_frames, args.height, args.width,
|
||||
)
|
||||
|
||||
# 3. Score each scenario. The metric reads file paths directly out
|
||||
# of the row dict (reference, reference_take2, masks), so we
|
||||
# just attach the generated video path and forward.
|
||||
evaluator = create_evaluator(metrics=["physics_iq"], num_gpus=args.num_gpus)
|
||||
|
||||
samples: list[dict] = []
|
||||
matched: list[dict] = []
|
||||
for row in rows:
|
||||
video_path = args.videos_dir / _expected_filename(row)
|
||||
if not video_path.is_file():
|
||||
print(f"[eval] missing {video_path}; skipping.")
|
||||
continue
|
||||
# The physics_iq metric accepts file paths via its polymorphic
|
||||
# input handling — no need to load the tensors here.
|
||||
samples.append({"video": str(video_path), **row})
|
||||
matched.append(row)
|
||||
|
||||
all_results = evaluator.evaluate(samples=samples)
|
||||
evaluator.shutdown()
|
||||
|
||||
# 4. Aggregate per the upstream scoring recipe.
|
||||
metric = get_metric("physics_iq")
|
||||
components = metric.aggregate_components(
|
||||
[r["physics_iq"] for r in all_results]
|
||||
)
|
||||
|
||||
print()
|
||||
print("=== Physics-IQ aggregate ===")
|
||||
for name, value in components.items():
|
||||
print(f" {name:24s} {value:.4f}")
|
||||
|
||||
detailed = [
|
||||
{
|
||||
"scenario": row["auxiliary_info"]["scenario_id"],
|
||||
"view": row["view"],
|
||||
"scenario_name": row["auxiliary_info"]["scenario_name"],
|
||||
"score": results["physics_iq"].score,
|
||||
}
|
||||
for row, results in zip(matched, all_results)
|
||||
]
|
||||
out = args.scores_out or (args.videos_dir / "scores.json")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(
|
||||
{"aggregate": components, "per_scenario": detailed},
|
||||
indent=2,
|
||||
))
|
||||
print(f"\n[done] per-scenario scores → {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,155 @@
|
||||
"""End-to-end VBench: dataset → generate → score → aggregate.
|
||||
|
||||
Iterates the VBench prompt corpus, generates one video per prompt with
|
||||
LTX2, scores each generated video against the requested ``vbench.*``
|
||||
sub-metrics, and prints per-metric averages over the run.
|
||||
|
||||
Re-running with ``--skip-generation`` reuses any mp4 already on disk
|
||||
under ``--videos-dir``, so you can iterate on metric selection without
|
||||
re-paying the generation cost.
|
||||
|
||||
Example — quick smoke run on 4 prompts from the ``aesthetic_quality``
|
||||
dimension across 2 GPUs::
|
||||
|
||||
python examples/inference/eval/bench_vbench.py \\
|
||||
--dimensions aesthetic_quality \\
|
||||
--limit 4 --num-gpus 2 \\
|
||||
--videos-dir outputs_video/vbench_smoke
|
||||
|
||||
Full benchmark on a single dimension::
|
||||
|
||||
python examples/inference/eval/bench_vbench.py \\
|
||||
--dimensions subject_consistency --num-gpus 8
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.datasets import get_dataset
|
||||
|
||||
|
||||
def _slugify(prompt: str, max_len: int = 100) -> str:
|
||||
"""Filesystem-safe filename stem; mirrors VBench's official convention."""
|
||||
s = re.sub(r'[\\/:*?"<>|]', "", prompt[:max_len]).strip().strip(".")
|
||||
return re.sub(r"\s+", " ", s) or "output"
|
||||
|
||||
|
||||
def _generate_videos(prompts: list[str], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
todo = [(p, videos_dir / f"{_slugify(p)}.mp4") for p in prompts]
|
||||
todo = [(p, out) for (p, out) in todo if not out.is_file()]
|
||||
if not todo:
|
||||
print(f"[gen] all {len(prompts)} videos already present; skipping.")
|
||||
return
|
||||
|
||||
print(f"[gen] {len(todo)}/{len(prompts)} prompts to render with {model} "
|
||||
f"({num_frames}x{height}x{width})...")
|
||||
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
|
||||
try:
|
||||
for prompt, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=prompt, output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
)
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dimensions", default="aesthetic_quality,subject_consistency",
|
||||
help="Comma-separated VBench dimensions (or 'all').")
|
||||
p.add_argument("--limit", type=int, default=None,
|
||||
help="Truncate to first N prompts for smoke runs.")
|
||||
p.add_argument("--videos-dir", type=Path,
|
||||
default=Path("outputs_video/bench_vbench"))
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the text→video generator to use.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--fps", type=float, default=24.0,
|
||||
help="Frame-rate annotation passed to fps-aware metrics.")
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Re-score existing videos under --videos-dir without "
|
||||
"regenerating.")
|
||||
p.add_argument("--scores-out", type=Path, default=None,
|
||||
help="Where to dump per-prompt scores as JSON. "
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
args = p.parse_args()
|
||||
|
||||
# 1. Pull prompts from VBench.
|
||||
dims_arg: list[str] | str = (
|
||||
args.dimensions if args.dimensions == "all"
|
||||
else [d.strip() for d in args.dimensions.split(",") if d.strip()]
|
||||
)
|
||||
ds = get_dataset("vbench", dimensions=dims_arg)
|
||||
rows = list(ds)[: args.limit]
|
||||
print(f"[load] VBench: {len(rows)} prompts across {ds.dimensions}")
|
||||
|
||||
# 2. Generate (or reuse) one mp4 per prompt.
|
||||
if not args.skip_generation:
|
||||
_generate_videos(
|
||||
[row["prompt"] for row in rows],
|
||||
args.videos_dir, args.model, args.num_gpus,
|
||||
args.num_frames, args.height, args.width,
|
||||
)
|
||||
|
||||
# 3. Score each video against the requested vbench sub-metrics.
|
||||
metric_names = sorted(set(f"vbench.{d}" for d in ds.dimensions))
|
||||
print(f"[eval] metrics: {metric_names}")
|
||||
evaluator = create_evaluator(metrics=metric_names, num_gpus=args.num_gpus)
|
||||
|
||||
samples: list[dict] = []
|
||||
matched_rows: list[dict] = []
|
||||
for row in rows:
|
||||
video_path = args.videos_dir / f"{_slugify(row['prompt'])}.mp4"
|
||||
if not video_path.is_file():
|
||||
print(f"[eval] missing {video_path}; skipping this row.")
|
||||
continue
|
||||
# Pass the path; the worker decodes lazily so memory stays bounded.
|
||||
samples.append({
|
||||
"video": str(video_path),
|
||||
"fps": args.fps,
|
||||
**row, # prompt / aux / dims
|
||||
})
|
||||
matched_rows.append(row)
|
||||
|
||||
all_results = evaluator.evaluate(samples=samples)
|
||||
evaluator.shutdown()
|
||||
|
||||
# 4. Aggregate per-metric.
|
||||
by_metric: dict[str, list[float]] = defaultdict(list)
|
||||
detailed: list[dict] = []
|
||||
for row, results in zip(matched_rows, all_results):
|
||||
scores = {name: r.score for name, r in results.items()}
|
||||
detailed.append({"prompt": row["prompt"], "scores": scores})
|
||||
for name, score in scores.items():
|
||||
if score is not None:
|
||||
by_metric[name].append(score)
|
||||
|
||||
print()
|
||||
print("=== per-metric averages ===")
|
||||
for name in sorted(by_metric):
|
||||
avg = sum(by_metric[name]) / len(by_metric[name])
|
||||
print(f" {name:42s} {avg:.4f} (n={len(by_metric[name])})")
|
||||
|
||||
out = args.scores_out or (args.videos_dir / "scores.json")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(detailed, indent=2))
|
||||
print(f"\n[done] per-prompt scores → {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,167 @@
|
||||
"""End-to-end: generate a video with LTX2 and score it with VBench.
|
||||
|
||||
Pipeline:
|
||||
prompt → LTX2-Base → mp4 → fastvideo.eval → vbench scores
|
||||
|
||||
Run::
|
||||
|
||||
pip install -e .[eval]
|
||||
git submodule update --init fastvideo/third_party/eval/vbench
|
||||
|
||||
python examples/inference/eval/eval_ltx2_vbench.py
|
||||
# or with 4 GPUs and the distilled checkpoint:
|
||||
python examples/inference/eval/eval_ltx2_vbench.py \
|
||||
--model FastVideo/LTX2-Distilled-Diffusers --num-gpus 4
|
||||
|
||||
The default metric set covers the vbench sub-metrics that are
|
||||
meaningful for an arbitrary text→video sample — i.e. those that need
|
||||
only the generated video (and optionally fps + the source prompt).
|
||||
Structured-prompt metrics like ``vbench.color``, ``vbench.scene``,
|
||||
``vbench.multiple_objects`` etc. are *not* on by default — they only
|
||||
make sense when the prompt is built to a specific schema, and they
|
||||
require GRiT/detectron2 setup. Pass them via ``--metrics`` if you have
|
||||
a matching prompt.
|
||||
|
||||
First-time runs download CLIP, DINO, RAFT, AMT, ViCLIP, and MUSIQ
|
||||
weights to ``~/.cache/fastvideo/eval/models/`` and
|
||||
``~/.cache/torch/hub/`` (~few GB total). Subsequent runs are fast.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
DEFAULT_METRICS = [
|
||||
# No-input metrics: just need the generated frames.
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.background_consistency", # DINO on background patches
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
# Need fps annotation:
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude
|
||||
# Need the source prompt:
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the LTX2 checkpoint.")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4",
|
||||
help="Where to save the generated mp4.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--prompt", default=PROMPT)
|
||||
p.add_argument("--fps", type=float, default=24.0,
|
||||
help="Frame-rate annotation passed to fps-aware metrics "
|
||||
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
|
||||
"by default.")
|
||||
p.add_argument("--metrics", default=",".join(DEFAULT_METRICS),
|
||||
help="Comma-separated metric names. Pass 'all' for every "
|
||||
"registered metric, or e.g. 'vbench' for the whole group.")
|
||||
p.add_argument("--scores-out", default="outputs_video/ltx2_eval/scores.json")
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Reuse an existing --output video instead of regenerating.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def generate(args: argparse.Namespace) -> Path:
|
||||
out = Path(args.output)
|
||||
if args.skip_generation and out.is_file():
|
||||
print(f"[gen] reusing existing video at {out}")
|
||||
return out
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
print(f"[gen] loading {args.model} ({args.num_gpus} GPU)...")
|
||||
generator = VideoGenerator.from_pretrained(args.model, num_gpus=args.num_gpus)
|
||||
try:
|
||||
print(f"[gen] generating to {out}...")
|
||||
generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
output_path=str(out),
|
||||
save_video=True,
|
||||
num_frames=args.num_frames,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
return out
|
||||
|
||||
|
||||
def evaluate_video(video_path: Path, prompt: str, fps: float,
|
||||
metric_names) -> dict:
|
||||
print(f"[eval] loading video from {video_path}...")
|
||||
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
|
||||
video = video.unsqueeze(0) # → (1, T, C, H, W)
|
||||
|
||||
print(f"[eval] building evaluator: {metric_names}")
|
||||
evaluator = create_evaluator(metrics=metric_names, device="cuda")
|
||||
|
||||
print(f"[eval] running ({video.shape[1]} frames @ {fps} fps)...")
|
||||
results = evaluator.evaluate(
|
||||
video=video,
|
||||
text_prompt=[prompt],
|
||||
fps=fps,
|
||||
)
|
||||
|
||||
if isinstance(results, list):
|
||||
results = results[0] # batch of 1
|
||||
|
||||
return {
|
||||
name: {"score": r.score, "details": r.details}
|
||||
for name, r in results.items()
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.metrics.strip() == "all":
|
||||
metric_names = "all"
|
||||
else:
|
||||
metric_names = [m.strip() for m in args.metrics.split(",") if m.strip()]
|
||||
|
||||
video_path = generate(args)
|
||||
scores = evaluate_video(video_path, args.prompt, args.fps, metric_names)
|
||||
|
||||
print("\n=== VBench scores ===")
|
||||
for name, payload in scores.items():
|
||||
print(f" {name}: {payload['score']}")
|
||||
|
||||
out = Path(args.scores_out)
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(
|
||||
{"video": str(video_path), "prompt": args.prompt, "scores": scores},
|
||||
indent=2,
|
||||
))
|
||||
print(f"[done] scores written to {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Score a folder of videos in parallel across multiple GPUs.
|
||||
|
||||
Uses :meth:`Evaluator.evaluate(samples=[...])`, which round-robins each
|
||||
sample dict across the GPU replicas the evaluator was built with.
|
||||
|
||||
Example::
|
||||
|
||||
python examples/inference/eval/score_folder.py \\
|
||||
--videos generated/ \\
|
||||
--metrics vbench.aesthetic_quality,vbench.subject_consistency \\
|
||||
--num-gpus 4 \\
|
||||
--output scores.json
|
||||
|
||||
Pair each generated video with a same-name reference video (e.g.
|
||||
``ref/<stem>.mp4``) by passing ``--reference-dir``::
|
||||
|
||||
python examples/inference/eval/score_folder.py \\
|
||||
--videos generated/ --reference-dir ref/ \\
|
||||
--metrics common.psnr,common.ssim,common.lpips \\
|
||||
--num-gpus 4
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
|
||||
def _list_videos(directory: Path) -> list[Path]:
|
||||
exts = {".mp4", ".avi", ".mov", ".mkv", ".gif"}
|
||||
return sorted(p for p in directory.iterdir() if p.suffix.lower() in exts)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--videos", type=Path, required=True,
|
||||
help="Directory of generated videos.")
|
||||
p.add_argument("--reference-dir", type=Path, default=None,
|
||||
help="Directory of reference videos with matching stems.")
|
||||
p.add_argument("--metrics", default="vbench.aesthetic_quality")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--fps", type=float, default=None,
|
||||
help="Frame-rate annotation for fps-aware metrics.")
|
||||
p.add_argument("--output", type=Path, default=Path("scores.json"))
|
||||
args = p.parse_args()
|
||||
|
||||
video_paths = _list_videos(args.videos)
|
||||
if not video_paths:
|
||||
raise SystemExit(f"No videos under {args.videos}")
|
||||
print(f"Found {len(video_paths)} videos in {args.videos}")
|
||||
|
||||
metrics: list[str] | str = (
|
||||
args.metrics if args.metrics == "all"
|
||||
else [m.strip() for m in args.metrics.split(",") if m.strip()]
|
||||
)
|
||||
evaluator = create_evaluator(metrics=metrics, num_gpus=args.num_gpus)
|
||||
|
||||
# Build per-video sample dicts holding *paths*, not pre-loaded
|
||||
# tensors. Each path is decoded inside the worker thread that picks
|
||||
# up its sample, so peak resident memory is bounded by num_gpus
|
||||
# rather than scaling with the size of the folder.
|
||||
samples: list[dict] = []
|
||||
for vp in video_paths:
|
||||
sample: dict = {"video": str(vp)}
|
||||
if args.reference_dir is not None:
|
||||
ref_path = args.reference_dir / vp.name
|
||||
if not ref_path.is_file():
|
||||
raise FileNotFoundError(f"Missing reference for {vp.name} at {ref_path}")
|
||||
sample["reference"] = str(ref_path)
|
||||
if args.fps is not None:
|
||||
sample["fps"] = args.fps
|
||||
samples.append(sample)
|
||||
|
||||
print(f"Scoring with {len(evaluator.metric_names)} metric(s) "
|
||||
f"on {evaluator.num_gpus} GPU(s)...")
|
||||
all_results = evaluator.evaluate(samples=samples)
|
||||
evaluator.shutdown()
|
||||
|
||||
payload = [
|
||||
{
|
||||
"video": str(vp),
|
||||
"scores": {name: r.score for name, r in results.items()},
|
||||
}
|
||||
for vp, results in zip(video_paths, all_results)
|
||||
]
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(json.dumps(payload, indent=2))
|
||||
print(f"Wrote {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Score one video on one GPU.
|
||||
|
||||
Smallest possible use of ``fastvideo.eval``: load an mp4, build an
|
||||
:class:`Evaluator` for the requested metric set, run it.
|
||||
|
||||
Examples::
|
||||
|
||||
# Reference-free (just the generated video):
|
||||
python examples/inference/eval/score_video.py \\
|
||||
--video clip.mp4 \\
|
||||
--metrics vbench.aesthetic_quality,vbench.imaging_quality
|
||||
|
||||
# Reference-paired (compare against ground truth):
|
||||
python examples/inference/eval/score_video.py \\
|
||||
--video gen.mp4 --reference ref.mp4 \\
|
||||
--metrics common.psnr,common.ssim,common.lpips
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--video", required=True, help="Path to the generated mp4.")
|
||||
p.add_argument("--reference", default=None,
|
||||
help="Optional path to a reference mp4 (for paired metrics).")
|
||||
p.add_argument("--metrics", default="common.psnr,common.ssim",
|
||||
help="Comma-separated metric names, or a group name like 'vbench'.")
|
||||
p.add_argument("--device", default="cuda:0")
|
||||
p.add_argument("--text-prompt", default=None,
|
||||
help="Text prompt for prompt-aware metrics "
|
||||
"(vbench.overall_consistency, etc.).")
|
||||
p.add_argument("--fps", type=float, default=None,
|
||||
help="Frame-rate annotation for fps-aware metrics "
|
||||
"(vbench.dynamic_degree, etc.).")
|
||||
args = p.parse_args()
|
||||
|
||||
metrics: list[str] | str = (
|
||||
args.metrics if args.metrics in ("all",)
|
||||
else [m.strip() for m in args.metrics.split(",") if m.strip()]
|
||||
)
|
||||
evaluator = create_evaluator(metrics=metrics, device=args.device)
|
||||
|
||||
sample: dict = {"video": load_video(args.video)}
|
||||
if args.reference is not None:
|
||||
sample["reference"] = load_video(args.reference)
|
||||
if args.text_prompt is not None:
|
||||
sample["text_prompt"] = args.text_prompt
|
||||
if args.fps is not None:
|
||||
sample["fps"] = args.fps
|
||||
|
||||
results = evaluator.evaluate(**sample)
|
||||
evaluator.shutdown()
|
||||
|
||||
print(json.dumps(
|
||||
{name: r.score for name, r in results.items()},
|
||||
indent=2,
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
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():
|
||||
|
||||
@@ -10,6 +10,35 @@ set -ex
|
||||
|
||||
echo "Building fastvideo-kernel..."
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Neutralise conda-injected compiler toolchains.
|
||||
#
|
||||
# Conda compiler packages (gcc_linux-aarch64, gxx_linux-64, etc.) set
|
||||
# CMAKE_ARGS, CFLAGS, CXXFLAGS, and LDFLAGS on activation. When multiple
|
||||
# toolchains are installed the variables can reference a *cross*-compiler
|
||||
# that doesn't match the host (e.g. aarch64-conda-linux-gnu-c++ on x86_64).
|
||||
# Even when the correct toolchain is active, the flags it injects
|
||||
# (-march=nocona, -mtune=haswell, …) can conflict with nvcc's host-compiler
|
||||
# expectations. Clear them so CMake discovers the system compiler instead.
|
||||
# ---------------------------------------------------------------------------
|
||||
if [[ -n "${CONDA_PREFIX:-}" ]]; then
|
||||
_need_clean=0
|
||||
# Detect conda cross-compiler that doesn't match the host.
|
||||
_host_arch="$(uname -m)"
|
||||
if [[ "${CXX:-}" == *"conda"* ]] || [[ "${CC:-}" == *"conda"* ]]; then
|
||||
_need_clean=1
|
||||
fi
|
||||
if [[ "${CMAKE_ARGS:-}" == *"conda"* ]]; then
|
||||
_need_clean=1
|
||||
fi
|
||||
if (( _need_clean )); then
|
||||
echo "NOTE: Clearing conda-injected compiler settings (CC/CXX/CMAKE_ARGS/CFLAGS/...)"
|
||||
echo " to use the system compiler for CUDA extension builds."
|
||||
unset CC CXX CMAKE_ARGS CFLAGS CXXFLAGS LDFLAGS
|
||||
fi
|
||||
unset _need_clean _host_arch
|
||||
fi
|
||||
|
||||
# Ensure submodules are initialized if needed (tk)
|
||||
git submodule update --init --recursive
|
||||
|
||||
@@ -32,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)
|
||||
|
||||
@@ -23,7 +23,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"torch>=2.5.0",
|
||||
"triton>=2.0.0",
|
||||
"triton>=2.0.0; sys_platform == 'linux'",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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__
|
||||
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ComponentConfig,
|
||||
ContinuationState,
|
||||
EngineConfig,
|
||||
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,
|
||||
load_raw_config,
|
||||
load_run_config,
|
||||
load_serve_config,
|
||||
parse_config,
|
||||
)
|
||||
from fastvideo.api.results import GenerationResult
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
"ComponentConfig",
|
||||
"ContinuationState",
|
||||
"ConfigValidationError",
|
||||
"EngineConfig",
|
||||
"GenerationResult",
|
||||
"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",
|
||||
"load_raw_config",
|
||||
"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",
|
||||
]
|
||||
@@ -0,0 +1,623 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from dataclasses import fields, is_dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
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_PATHS_ATTR,
|
||||
bind_generation_request_raw,
|
||||
get_explicit_paths,
|
||||
reset_tracking_roots,
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
RequestRuntimeConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
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)}
|
||||
_SAMPLING_FIELD_NAMES = {field.name for field in fields(SamplingConfig)}
|
||||
_RUNTIME_FIELD_NAMES = {field.name for field in fields(RequestRuntimeConfig)}
|
||||
_OUTPUT_FIELD_NAMES = {field.name for field in fields(OutputConfig)}
|
||||
_MISSING = object()
|
||||
_LEGACY_REQUEST_ALIASES = {
|
||||
"neg_prompt": "negative_prompt",
|
||||
}
|
||||
_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:
|
||||
if isinstance(config, GeneratorConfig):
|
||||
return config
|
||||
return parse_config(GeneratorConfig, config)
|
||||
|
||||
|
||||
def load_generator_config_from_file(
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> GeneratorConfig:
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = normalize_overrides(overrides)
|
||||
|
||||
if _looks_like_run_or_serve_config(raw):
|
||||
if normalized_overrides:
|
||||
raw = apply_overrides(raw, normalized_overrides)
|
||||
return parse_config(GeneratorConfig, raw["generator"])
|
||||
|
||||
if normalized_overrides:
|
||||
adjusted = normalized_overrides
|
||||
if all(key.startswith("generator.") for key in adjusted):
|
||||
adjusted = {key[len("generator."):]: value for key, value in adjusted.items()}
|
||||
raw = apply_overrides(raw, adjusted)
|
||||
|
||||
return parse_config(GeneratorConfig, raw)
|
||||
|
||||
|
||||
def legacy_from_pretrained_to_config(
|
||||
model_path: str,
|
||||
kwargs: Mapping[str, Any],
|
||||
) -> GeneratorConfig:
|
||||
raw: dict[str, Any] = {"model_path": model_path}
|
||||
engine: dict[str, Any] = {}
|
||||
parallelism: dict[str, Any] = {}
|
||||
offload: dict[str, Any] = {}
|
||||
compile_config: dict[str, Any] = {}
|
||||
pipeline: dict[str, Any] = {}
|
||||
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":
|
||||
raw["revision"] = value
|
||||
elif key == "trust_remote_code":
|
||||
raw["trust_remote_code"] = value
|
||||
elif key == "num_gpus":
|
||||
engine["num_gpus"] = value
|
||||
elif key == "distributed_executor_backend":
|
||||
engine["execution_backend"] = value
|
||||
elif key in {"tp_size", "sp_size", "hsdp_replicate_dim", "hsdp_shard_dim", "dist_timeout"}:
|
||||
parallelism[key] = value
|
||||
elif key == "dit_cpu_offload":
|
||||
offload["dit"] = value
|
||||
elif key == "dit_layerwise_offload":
|
||||
offload["dit_layerwise"] = value
|
||||
elif key == "text_encoder_cpu_offload":
|
||||
offload["text_encoder"] = value
|
||||
elif key == "image_encoder_cpu_offload":
|
||||
offload["image_encoder"] = value
|
||||
elif key == "vae_cpu_offload":
|
||||
offload["vae"] = value
|
||||
elif key == "pin_cpu_memory":
|
||||
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":
|
||||
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":
|
||||
quantization["text_encoder_quant"] = value
|
||||
elif key == "workload_type":
|
||||
pipeline["workload_type"] = value
|
||||
elif key == "lora_path":
|
||||
components["lora_path"] = value
|
||||
elif key == "override_pipeline_cls_name":
|
||||
components["override_pipeline_cls_name"] = value
|
||||
elif key == "override_transformer_cls_name":
|
||||
components["override_transformer_cls_name"] = value
|
||||
elif key == "pipeline_config":
|
||||
if isinstance(value, str):
|
||||
components["pipeline_config_path"] = value
|
||||
else:
|
||||
experimental[key] = deepcopy(value)
|
||||
elif key == "override_text_encoder_safetensors":
|
||||
components["text_encoder_weights"] = value
|
||||
elif key == "init_weights_from_safetensors":
|
||||
components["transformer_weights"] = value
|
||||
elif key == "init_weights_from_safetensors_2":
|
||||
components["transformer_2_weights"] = value
|
||||
else:
|
||||
experimental[key] = deepcopy(value)
|
||||
|
||||
if parallelism:
|
||||
engine["parallelism"] = parallelism
|
||||
if offload:
|
||||
engine["offload"] = offload
|
||||
if compile_config:
|
||||
engine["compile"] = compile_config
|
||||
if quantization:
|
||||
engine["quantization"] = quantization
|
||||
if engine:
|
||||
raw["engine"] = engine
|
||||
|
||||
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:
|
||||
raw["pipeline"] = pipeline
|
||||
|
||||
return parse_config(GeneratorConfig, raw)
|
||||
|
||||
|
||||
def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, Any], ) -> FastVideoArgs:
|
||||
normalized = normalize_generator_config(config)
|
||||
unsupported = []
|
||||
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 unsupported:
|
||||
joined = ", ".join(unsupported)
|
||||
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
|
||||
|
||||
engine = normalized.engine
|
||||
kwargs: dict[str, Any] = {
|
||||
"model_path": normalized.model_path,
|
||||
"revision": normalized.revision,
|
||||
"trust_remote_code": normalized.trust_remote_code,
|
||||
"num_gpus": engine.num_gpus,
|
||||
"distributed_executor_backend": engine.execution_backend,
|
||||
"tp_size": engine.parallelism.tp_size,
|
||||
"sp_size": engine.parallelism.sp_size,
|
||||
"hsdp_replicate_dim": engine.parallelism.hsdp_replicate_dim,
|
||||
"hsdp_shard_dim": engine.parallelism.hsdp_shard_dim,
|
||||
"dist_timeout": engine.parallelism.dist_timeout,
|
||||
"dit_cpu_offload": engine.offload.dit,
|
||||
"dit_layerwise_offload": engine.offload.dit_layerwise,
|
||||
"text_encoder_cpu_offload": engine.offload.text_encoder,
|
||||
"image_encoder_cpu_offload": engine.offload.image_encoder,
|
||||
"vae_cpu_offload": engine.offload.vae,
|
||||
"pin_cpu_memory": engine.offload.pin_cpu_memory,
|
||||
"enable_torch_compile": engine.compile.enabled,
|
||||
"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:
|
||||
kwargs["override_text_encoder_quant"] = quantization.text_encoder_quant
|
||||
if quantization is not None and quantization.transformer_quant is not None:
|
||||
kwargs["transformer_quant"] = quantization.transformer_quant
|
||||
|
||||
components = normalized.pipeline.components
|
||||
if components.pipeline_config_path is not None:
|
||||
kwargs["pipeline_config"] = components.pipeline_config_path
|
||||
if components.lora_path is not None:
|
||||
kwargs["lora_path"] = components.lora_path
|
||||
if components.override_pipeline_cls_name is not None:
|
||||
kwargs["override_pipeline_cls_name"] = components.override_pipeline_cls_name
|
||||
if components.override_transformer_cls_name is not None:
|
||||
kwargs["override_transformer_cls_name"] = components.override_transformer_cls_name
|
||||
if components.text_encoder_weights is not None:
|
||||
kwargs["override_text_encoder_safetensors"] = components.text_encoder_weights
|
||||
if components.transformer_weights is not None:
|
||||
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
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def normalize_generation_request(request: GenerationRequest | Mapping[str, Any], ) -> GenerationRequest:
|
||||
normalized = (request if isinstance(request, GenerationRequest) else parse_config(GenerationRequest, request))
|
||||
|
||||
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
|
||||
|
||||
|
||||
def legacy_generate_call_to_request(
|
||||
prompt: str | None,
|
||||
sampling_param: SamplingParam | None,
|
||||
*,
|
||||
mouse_cond: Any | None = None,
|
||||
keyboard_cond: Any | None = None,
|
||||
grid_sizes: Any | None = None,
|
||||
legacy_kwargs: Mapping[str, Any] | None = None,
|
||||
) -> GenerationRequest:
|
||||
raw = _sampling_param_to_request_raw(sampling_param)
|
||||
if prompt is not None:
|
||||
raw["prompt"] = prompt
|
||||
|
||||
for key, value in (legacy_kwargs or {}).items():
|
||||
_apply_request_field(raw, key, value)
|
||||
|
||||
if mouse_cond is not None:
|
||||
raw.setdefault("inputs", {})["mouse_cond"] = mouse_cond
|
||||
if keyboard_cond is not None:
|
||||
raw.setdefault("inputs", {})["keyboard_cond"] = keyboard_cond
|
||||
if grid_sizes is not None:
|
||||
raw.setdefault("inputs", {})["grid_sizes"] = grid_sizes
|
||||
|
||||
normalized = parse_config(GenerationRequest, raw)
|
||||
bind_generation_request_raw(normalized, raw)
|
||||
return normalized
|
||||
|
||||
|
||||
def request_to_sampling_param(
|
||||
request: GenerationRequest,
|
||||
*,
|
||||
model_path: str,
|
||||
) -> SamplingParam:
|
||||
if request.plan is not None:
|
||||
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
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:
|
||||
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}")
|
||||
|
||||
sampling_param.__post_init__()
|
||||
sampling_param.check_sampling_param()
|
||||
return sampling_param
|
||||
|
||||
|
||||
def expand_request_prompt_batch(request: GenerationRequest, ) -> list[GenerationRequest]:
|
||||
if not isinstance(request.prompt, list):
|
||||
return [request]
|
||||
|
||||
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)
|
||||
requests.append(single_request)
|
||||
return requests
|
||||
|
||||
|
||||
def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
|
||||
return isinstance(raw.get("generator"), Mapping)
|
||||
|
||||
|
||||
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]:
|
||||
if sampling_param is None:
|
||||
return {}
|
||||
|
||||
raw: dict[str, Any] = {}
|
||||
for key, value in shallow_asdict(sampling_param).items():
|
||||
if key == "prompt":
|
||||
continue
|
||||
_apply_request_field(raw, key, deepcopy(value))
|
||||
return raw
|
||||
|
||||
|
||||
def _apply_request_field(
|
||||
raw: dict[str, Any],
|
||||
key: str,
|
||||
value: Any,
|
||||
) -> None:
|
||||
key = _LEGACY_REQUEST_ALIASES.get(key, key)
|
||||
if key == "negative_prompt":
|
||||
raw["negative_prompt"] = value
|
||||
return
|
||||
if key in _INPUT_FIELD_NAMES:
|
||||
raw.setdefault("inputs", {})[key] = value
|
||||
return
|
||||
if key in _SAMPLING_FIELD_NAMES:
|
||||
raw.setdefault("sampling", {})[key] = value
|
||||
return
|
||||
if key in _RUNTIME_FIELD_NAMES:
|
||||
raw.setdefault("runtime", {})[key] = value
|
||||
return
|
||||
if key in _OUTPUT_FIELD_NAMES:
|
||||
raw.setdefault("output", {})[key] = value
|
||||
return
|
||||
raw.setdefault("extensions", {})[key] = value
|
||||
|
||||
|
||||
def request_to_pipeline_overrides(request: GenerationRequest) -> dict[str, Any]:
|
||||
overrides: dict[str, Any] = {}
|
||||
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]:
|
||||
"""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:
|
||||
updates["negative_prompt"] = deepcopy(raw["negative_prompt"])
|
||||
|
||||
for section_name in ("inputs", "sampling", "runtime", "output"):
|
||||
section = raw.get(section_name)
|
||||
if not isinstance(section, Mapping):
|
||||
continue
|
||||
for key, value in section.items():
|
||||
updates[key] = deepcopy(value)
|
||||
|
||||
stage_overrides = raw.get("stage_overrides")
|
||||
if stage_overrides:
|
||||
updates.update(_flatten_stage_overrides(stage_overrides))
|
||||
|
||||
extensions = raw.get("extensions")
|
||||
if isinstance(extensions, Mapping):
|
||||
for key, value in extensions.items():
|
||||
updates[key] = deepcopy(value)
|
||||
|
||||
return updates
|
||||
|
||||
|
||||
def _flatten_stage_overrides(stage_overrides: Any) -> dict[str, Any]:
|
||||
if not isinstance(stage_overrides, Mapping):
|
||||
raise ValueError("GenerationRequest.stage_overrides must be a mapping")
|
||||
|
||||
flattened: dict[str, Any] = {}
|
||||
for stage_name, overrides in stage_overrides.items():
|
||||
if not isinstance(overrides, Mapping):
|
||||
raise ValueError(f"GenerationRequest.stage_overrides.{stage_name} must be a mapping")
|
||||
for key, value in overrides.items():
|
||||
if key in flattened and flattened[key] != value:
|
||||
raise ValueError(f"Conflicting stage override for {key!r} across stages")
|
||||
flattened[key] = deepcopy(value)
|
||||
return flattened
|
||||
|
||||
|
||||
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,
|
||||
field_name: str,
|
||||
index: int,
|
||||
) -> None:
|
||||
value = getattr(source_request.inputs, field_name)
|
||||
if not isinstance(value, list):
|
||||
return
|
||||
_validate_batched_input_length(source_request.prompt, value, field_name)
|
||||
setattr(target_request.inputs, field_name, deepcopy(value[index]))
|
||||
|
||||
|
||||
def _validate_batched_input_length(
|
||||
prompts: str | list[str] | None,
|
||||
values: list[Any],
|
||||
field_name: str,
|
||||
) -> None:
|
||||
if not isinstance(prompts, list):
|
||||
return
|
||||
if len(values) != len(prompts):
|
||||
raise ValueError(f"GenerationRequest.inputs.{field_name} must have the same length as request.prompt")
|
||||
|
||||
|
||||
__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",
|
||||
]
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class ConfigValidationError(ValueError):
|
||||
"""Validation error that keeps track of the nested config path."""
|
||||
|
||||
def __init__(self, path: str, message: str):
|
||||
self.path = path
|
||||
self.message = message
|
||||
super().__init__(str(self))
|
||||
|
||||
def __str__(self) -> str:
|
||||
if self.path:
|
||||
return f"{self.path}: {self.message}"
|
||||
return self.message
|
||||
@@ -0,0 +1,110 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
from collections.abc import Mapping
|
||||
|
||||
import yaml
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
|
||||
|
||||
def parse_cli_overrides(overrides: list[str]) -> dict[str, Any]:
|
||||
"""Parse ``--dotted.key value`` style overrides into a flat mapping."""
|
||||
parsed: dict[str, Any] = {}
|
||||
index = 0
|
||||
while index < len(overrides):
|
||||
token = overrides[index]
|
||||
if not token.startswith("--"):
|
||||
raise ValueError(f"Expected --dotted.key, got {token!r}")
|
||||
|
||||
key = token[2:]
|
||||
if not key:
|
||||
raise ValueError("Override key cannot be empty")
|
||||
|
||||
if "=" in key:
|
||||
key, raw_value = key.split("=", 1)
|
||||
else:
|
||||
index += 1
|
||||
if index >= len(overrides):
|
||||
raise ValueError(f"Missing value for override {token!r}")
|
||||
raw_value = overrides[index]
|
||||
|
||||
parsed[_normalize_override_key(key)] = _cast_override_value(raw_value)
|
||||
index += 1
|
||||
|
||||
return parsed
|
||||
|
||||
|
||||
def apply_overrides(config: Mapping[str, Any], overrides: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Return a copy of ``config`` with dotted-key overrides applied."""
|
||||
merged = deepcopy(dict(config))
|
||||
for dotted_key, value in overrides.items():
|
||||
_apply_single_override(merged, dotted_key, value)
|
||||
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):
|
||||
raise ValueError(f"Invalid override path {dotted_key!r}")
|
||||
|
||||
cursor = config
|
||||
for depth, part in enumerate(parts[:-1]):
|
||||
existing = cursor.get(part)
|
||||
if existing is None:
|
||||
existing = {}
|
||||
cursor[part] = existing
|
||||
elif not isinstance(existing, dict):
|
||||
raise ConfigValidationError(
|
||||
".".join(parts[:depth + 1]),
|
||||
"cannot apply nested override through a non-mapping value",
|
||||
)
|
||||
cursor = existing
|
||||
|
||||
cursor[parts[-1]] = value
|
||||
|
||||
|
||||
def _cast_override_value(raw: str) -> Any:
|
||||
lowered = raw.lower()
|
||||
if lowered == "true":
|
||||
return True
|
||||
if lowered == "false":
|
||||
return False
|
||||
if lowered in {"none", "null"}:
|
||||
return None
|
||||
|
||||
try:
|
||||
return int(raw)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
try:
|
||||
return float(raw)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if raw.startswith("[") or raw.startswith("{"):
|
||||
try:
|
||||
return yaml.safe_load(raw)
|
||||
except yaml.YAMLError:
|
||||
pass
|
||||
|
||||
return raw
|
||||
|
||||
|
||||
def _normalize_override_key(key: str) -> str:
|
||||
return key.replace("-", "_")
|
||||
|
||||
|
||||
__all__ = ["apply_overrides", "normalize_overrides", "parse_cli_overrides"]
|
||||
@@ -0,0 +1,324 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import json
|
||||
import types
|
||||
from pathlib import Path
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Literal, TypeVar, Union, get_args, get_origin, get_type_hints
|
||||
|
||||
import yaml
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, normalize_overrides
|
||||
from fastvideo.api.request_metadata import (
|
||||
bind_generation_request_raw,
|
||||
bind_run_config_raw,
|
||||
bind_serve_config_raw,
|
||||
)
|
||||
from fastvideo.api.schema import GenerationRequest, RunConfig, ServeConfig
|
||||
|
||||
T = TypeVar("T")
|
||||
_UNION_ORIGINS = {types.UnionType, Union}
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class _DataclassSpec:
|
||||
cls: type[Any]
|
||||
type_hints: dict[str, Any]
|
||||
fields_by_name: dict[str, dataclasses.Field[Any]]
|
||||
|
||||
|
||||
def parse_config(config_type: type[T], raw: Mapping[str, Any] | T) -> T:
|
||||
"""Parse a nested mapping into a typed inference config object."""
|
||||
if isinstance(raw, config_type):
|
||||
return raw
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError("", f"expected mapping for {config_type.__name__}")
|
||||
parsed = _SchemaParser().parse_dataclass(config_type, raw, "")
|
||||
if config_type is GenerationRequest:
|
||||
return bind_generation_request_raw(parsed, raw)
|
||||
if config_type is RunConfig:
|
||||
return bind_run_config_raw(parsed, raw)
|
||||
if config_type is ServeConfig:
|
||||
return bind_serve_config_raw(parsed, raw)
|
||||
return parsed
|
||||
|
||||
|
||||
def config_to_dict(config: Any) -> Any:
|
||||
"""Serialize a typed config object into plain Python containers."""
|
||||
if dataclasses.is_dataclass(config) and not isinstance(config, type):
|
||||
return {field.name: config_to_dict(getattr(config, field.name)) for field in dataclasses.fields(config)}
|
||||
if isinstance(config, list):
|
||||
return [config_to_dict(item) for item in config]
|
||||
if isinstance(config, dict):
|
||||
return {key: config_to_dict(value) for key, value in config.items()}
|
||||
return config
|
||||
|
||||
|
||||
def load_config(
|
||||
config_type: type[T],
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> T:
|
||||
"""Load a typed config object from YAML or JSON."""
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = normalize_overrides(overrides)
|
||||
if normalized_overrides:
|
||||
raw = apply_overrides(raw, normalized_overrides)
|
||||
return parse_config(config_type, raw)
|
||||
|
||||
|
||||
def load_run_config(
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> RunConfig:
|
||||
return load_config(RunConfig, path, overrides)
|
||||
|
||||
|
||||
def load_serve_config(
|
||||
path: str | Path,
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> ServeConfig:
|
||||
return load_config(ServeConfig, path, overrides)
|
||||
|
||||
|
||||
def load_raw_config(path: str | Path) -> dict[str, Any]:
|
||||
config_path = Path(path)
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||
|
||||
with config_path.open(encoding="utf-8") as handle:
|
||||
raw = _load_raw_mapping(handle, config_path)
|
||||
|
||||
if raw is None:
|
||||
return {}
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError("", f"{config_path} must contain a top-level mapping")
|
||||
return dict(raw)
|
||||
|
||||
|
||||
def _load_raw_mapping(handle: Any, config_path: Path) -> Any:
|
||||
suffix = config_path.suffix.lower()
|
||||
if suffix in {".yaml", ".yml"}:
|
||||
return yaml.safe_load(handle)
|
||||
if suffix == ".json":
|
||||
return json.load(handle)
|
||||
raise ValueError(f"Unsupported config file format: {config_path}")
|
||||
|
||||
|
||||
class _SchemaParser:
|
||||
|
||||
def parse_dataclass(
|
||||
self,
|
||||
config_type: type[T],
|
||||
raw: Mapping[str, Any],
|
||||
path: str,
|
||||
) -> T:
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError(path, f"expected mapping for {config_type.__name__}")
|
||||
|
||||
spec = _get_dataclass_spec(config_type)
|
||||
self._validate_keys(raw, spec, path)
|
||||
|
||||
values: dict[str, Any] = {}
|
||||
for name, field in spec.fields_by_name.items():
|
||||
field_path = _join_path(path, name)
|
||||
if name in raw:
|
||||
values[name] = self.parse_value(spec.type_hints[name], raw[name], field_path)
|
||||
continue
|
||||
if _field_is_required(field):
|
||||
raise ConfigValidationError(field_path, "missing required field")
|
||||
|
||||
return config_type(**values)
|
||||
|
||||
def parse_value(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
if annotation is Any:
|
||||
return value
|
||||
|
||||
origin = get_origin(annotation)
|
||||
if origin in _UNION_ORIGINS:
|
||||
return self._parse_union(annotation, value, path)
|
||||
if origin is Literal:
|
||||
return self._parse_literal(annotation, value, path)
|
||||
if origin is list:
|
||||
return self._parse_list(annotation, value, path)
|
||||
if origin is dict:
|
||||
return self._parse_dict(annotation, value, path)
|
||||
if origin is tuple:
|
||||
return self._parse_tuple(annotation, value, path)
|
||||
if isinstance(annotation, type) and dataclasses.is_dataclass(annotation):
|
||||
return self.parse_dataclass(annotation, value, path)
|
||||
|
||||
scalar_parser = _SCALAR_PARSERS.get(annotation)
|
||||
if scalar_parser is not None:
|
||||
return scalar_parser(value, path)
|
||||
|
||||
return self._parse_instance(annotation, value, path)
|
||||
|
||||
def _validate_keys(
|
||||
self,
|
||||
raw: Mapping[str, Any],
|
||||
spec: _DataclassSpec,
|
||||
path: str,
|
||||
) -> None:
|
||||
for key in raw:
|
||||
if not isinstance(key, str):
|
||||
raise ConfigValidationError(path, "expected mapping keys to be strings")
|
||||
if key not in spec.fields_by_name:
|
||||
raise ConfigValidationError(_join_path(path, key), "unknown field")
|
||||
|
||||
def _parse_union(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
candidates = [candidate for candidate in get_args(annotation) if candidate is not type(None)]
|
||||
if value is None and len(candidates) != len(get_args(annotation)):
|
||||
return None
|
||||
if len(candidates) == 1:
|
||||
return self.parse_value(candidates[0], value, path)
|
||||
|
||||
errors: list[str] = []
|
||||
for candidate in candidates:
|
||||
try:
|
||||
return self.parse_value(candidate, value, path)
|
||||
except ConfigValidationError as exc:
|
||||
errors.append(exc.message)
|
||||
|
||||
expected = ", ".join(_type_name(candidate) for candidate in candidates)
|
||||
detail = errors[0] if errors else f"expected one of ({expected})"
|
||||
raise ConfigValidationError(path, detail)
|
||||
|
||||
def _parse_literal(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
allowed = get_args(annotation)
|
||||
if value not in allowed:
|
||||
raise ConfigValidationError(path, f"expected one of {sorted(allowed)!r}")
|
||||
return value
|
||||
|
||||
def _parse_list(self, annotation: Any, value: Any, path: str) -> list[Any]:
|
||||
if not isinstance(value, list):
|
||||
raise ConfigValidationError(path, "expected list")
|
||||
item_type = get_args(annotation)[0] if get_args(annotation) else Any
|
||||
return [self.parse_value(item_type, item, f"{path}[{index}]") for index, item in enumerate(value)]
|
||||
|
||||
def _parse_dict(self, annotation: Any, value: Any, path: str) -> dict[Any, Any]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ConfigValidationError(path, "expected mapping")
|
||||
|
||||
key_type, value_type = (get_args(annotation) + (Any, Any))[:2]
|
||||
parsed: dict[Any, Any] = {}
|
||||
for key, item in value.items():
|
||||
parsed_key = self._parse_dict_key(key_type, key, path)
|
||||
item_path = _join_path(path, str(key))
|
||||
parsed[parsed_key] = self.parse_value(value_type, item, item_path)
|
||||
return parsed
|
||||
|
||||
def _parse_tuple(self, annotation: Any, value: Any, path: str) -> tuple[Any, ...]:
|
||||
if not isinstance(value, list | tuple):
|
||||
raise ConfigValidationError(path, "expected tuple")
|
||||
|
||||
item_types = get_args(annotation)
|
||||
if len(item_types) == 2 and item_types[1] is Ellipsis:
|
||||
return tuple(self.parse_value(item_types[0], item, f"{path}[{index}]") for index, item in enumerate(value))
|
||||
|
||||
if len(value) != len(item_types):
|
||||
raise ConfigValidationError(path, f"expected tuple of length {len(item_types)}")
|
||||
|
||||
return tuple(
|
||||
self.parse_value(item_type, item, f"{path}[{index}]")
|
||||
for index, (item_type, item) in enumerate(zip(item_types, value, strict=True)))
|
||||
|
||||
def _parse_dict_key(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
if annotation is Any:
|
||||
return value
|
||||
if annotation is str:
|
||||
if not isinstance(value, str):
|
||||
raise ConfigValidationError(path, "expected string dictionary keys")
|
||||
return value
|
||||
if annotation is int:
|
||||
if not isinstance(value, int) or isinstance(value, bool):
|
||||
raise ConfigValidationError(path, "expected integer dictionary keys")
|
||||
return value
|
||||
return value
|
||||
|
||||
def _parse_instance(self, annotation: Any, value: Any, path: str) -> Any:
|
||||
if isinstance(annotation, type) and not isinstance(value, annotation):
|
||||
raise ConfigValidationError(path, f"expected {annotation.__name__}")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_bool(value: Any, path: str) -> bool:
|
||||
if type(value) is not bool:
|
||||
raise ConfigValidationError(path, "expected bool")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_int(value: Any, path: str) -> int:
|
||||
if not isinstance(value, int) or isinstance(value, bool):
|
||||
raise ConfigValidationError(path, "expected int")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_float(value: Any, path: str) -> float:
|
||||
if not isinstance(value, int | float) or isinstance(value, bool):
|
||||
raise ConfigValidationError(path, "expected float")
|
||||
return float(value)
|
||||
|
||||
|
||||
def _parse_str(value: Any, path: str) -> str:
|
||||
if not isinstance(value, str):
|
||||
raise ConfigValidationError(path, "expected str")
|
||||
return value
|
||||
|
||||
|
||||
_SCALAR_PARSERS: dict[Any, Any] = {
|
||||
bool: _parse_bool,
|
||||
int: _parse_int,
|
||||
float: _parse_float,
|
||||
str: _parse_str,
|
||||
}
|
||||
|
||||
|
||||
def _field_is_required(field: dataclasses.Field[Any]) -> bool:
|
||||
return (field.default is dataclasses.MISSING and field.default_factory is dataclasses.MISSING)
|
||||
|
||||
|
||||
def _get_dataclass_spec(config_type: type[Any]) -> _DataclassSpec:
|
||||
spec = _DATACLASS_SPEC_CACHE.get(config_type)
|
||||
if spec is not None:
|
||||
return spec
|
||||
|
||||
spec = _DataclassSpec(
|
||||
cls=config_type,
|
||||
type_hints=get_type_hints(config_type),
|
||||
fields_by_name={field.name: field
|
||||
for field in dataclasses.fields(config_type)},
|
||||
)
|
||||
_DATACLASS_SPEC_CACHE[config_type] = spec
|
||||
return spec
|
||||
|
||||
|
||||
_DATACLASS_SPEC_CACHE: dict[type[Any], _DataclassSpec] = {}
|
||||
|
||||
|
||||
def _join_path(prefix: str, suffix: str) -> str:
|
||||
if not prefix:
|
||||
return suffix
|
||||
return f"{prefix}.{suffix}"
|
||||
|
||||
|
||||
def _type_name(annotation: Any) -> str:
|
||||
origin = get_origin(annotation)
|
||||
if origin is not None:
|
||||
return str(annotation)
|
||||
if hasattr(annotation, "__name__"):
|
||||
return annotation.__name__
|
||||
return str(annotation)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"config_to_dict",
|
||||
"load_config",
|
||||
"load_raw_config",
|
||||
"load_run_config",
|
||||
"load_serve_config",
|
||||
"parse_config",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -0,0 +1,233 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Track which GenerationRequest fields the user explicitly provided.
|
||||
|
||||
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 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 Callable, Mapping
|
||||
import dataclasses
|
||||
from typing import Any, cast
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
PlannedStage,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
|
||||
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"
|
||||
|
||||
_TRACKED_REQUEST_TYPES = (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
SamplingConfig,
|
||||
RequestRuntimeConfig,
|
||||
OutputConfig,
|
||||
ContinuationState,
|
||||
PlannedStage,
|
||||
GenerationPlan,
|
||||
)
|
||||
|
||||
|
||||
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 recording while we walk the tree to install roots.
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, None)
|
||||
_set_tracking_roots(request, request, "")
|
||||
paths: set[str] = set()
|
||||
_record_value_paths(raw or {}, "", paths)
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, paths)
|
||||
return request
|
||||
|
||||
|
||||
def bind_run_config_raw(
|
||||
config: RunConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> RunConfig:
|
||||
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
|
||||
|
||||
|
||||
def bind_serve_config_raw(
|
||||
config: ServeConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> ServeConfig:
|
||||
default_request_raw = raw.get("default_request")
|
||||
if isinstance(default_request_raw, Mapping):
|
||||
bind_generation_request_raw(config.default_request, default_request_raw)
|
||||
else:
|
||||
bind_generation_request_raw(config.default_request, {})
|
||||
return config
|
||||
|
||||
|
||||
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()
|
||||
|
||||
|
||||
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, "")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Path recording
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _record_value_paths(
|
||||
value: Any,
|
||||
prefix: str,
|
||||
out: set[str],
|
||||
) -> None:
|
||||
"""Add every leaf path under *value* to *out*.
|
||||
|
||||
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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ensure_request_tracking() -> None:
|
||||
for config_type in _TRACKED_REQUEST_TYPES:
|
||||
_patch_tracking_setattr(config_type)
|
||||
|
||||
|
||||
def _patch_tracking_setattr(config_type: type[Any]) -> None:
|
||||
if getattr(config_type, _TRACKING_PATCHED_ATTR, False):
|
||||
return
|
||||
|
||||
original_setattr = cast(
|
||||
Callable[[Any, str, Any], None],
|
||||
config_type.__setattr__,
|
||||
)
|
||||
field_names = {field.name for field in dataclasses.fields(config_type)}
|
||||
|
||||
def _tracking_setattr(self: Any, name: str, value: Any) -> None:
|
||||
if name.startswith("_fastvideo_") or name not in field_names:
|
||||
original_setattr(self, name, value)
|
||||
return
|
||||
|
||||
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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tree walk to set tracking root/path on nested dataclasses
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _set_tracking_roots(
|
||||
root: GenerationRequest,
|
||||
obj: Any,
|
||||
prefix: str,
|
||||
) -> None:
|
||||
if not dataclasses.is_dataclass(obj) or isinstance(obj, type):
|
||||
return
|
||||
object.__setattr__(obj, _TRACKING_ROOT_ATTR, root)
|
||||
object.__setattr__(obj, _TRACKING_PATH_ATTR, prefix)
|
||||
for field in dataclasses.fields(obj):
|
||||
child = getattr(obj, field.name)
|
||||
child_path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
if dataclasses.is_dataclass(child) and not isinstance(child, type):
|
||||
_set_tracking_roots(root, child, child_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EXPLICIT_PATHS_ATTR",
|
||||
"bind_generation_request_raw",
|
||||
"bind_run_config_raw",
|
||||
"bind_serve_config_raw",
|
||||
"get_explicit_paths",
|
||||
"reset_tracking_roots",
|
||||
]
|
||||
@@ -0,0 +1,101 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from collections.abc import Mapping
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationResult:
|
||||
prompt: str | None = None
|
||||
prompt_index: int | None = None
|
||||
samples: Any | None = None
|
||||
frames: Any | None = None
|
||||
audio: Any | None = None
|
||||
size: tuple[int, int, int] | None = None
|
||||
generation_time: float | None = None
|
||||
logging_info: Any | None = None
|
||||
trajectory: Any | None = None
|
||||
trajectory_timesteps: Any | None = None
|
||||
trajectory_decoded: Any | None = None
|
||||
video_path: str | None = None
|
||||
peak_memory_mb: float | None = None
|
||||
state: ContinuationState | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_legacy_result(
|
||||
cls,
|
||||
result: Mapping[str, Any],
|
||||
) -> GenerationResult:
|
||||
prompt = result.get("prompt")
|
||||
if prompt is None:
|
||||
prompt = result.get("prompts")
|
||||
|
||||
extra = {
|
||||
key: value
|
||||
for key, value in result.items() if key not in {
|
||||
"prompt",
|
||||
"prompt_index",
|
||||
"prompts",
|
||||
"samples",
|
||||
"frames",
|
||||
"audio",
|
||||
"size",
|
||||
"generation_time",
|
||||
"logging_info",
|
||||
"trajectory",
|
||||
"trajectory_timesteps",
|
||||
"trajectory_decoded",
|
||||
"video_path",
|
||||
"peak_memory_mb",
|
||||
"state",
|
||||
}
|
||||
}
|
||||
|
||||
return cls(
|
||||
prompt=prompt,
|
||||
prompt_index=result.get("prompt_index"),
|
||||
samples=result.get("samples"),
|
||||
frames=result.get("frames"),
|
||||
audio=result.get("audio"),
|
||||
size=result.get("size"),
|
||||
generation_time=result.get("generation_time"),
|
||||
logging_info=result.get("logging_info"),
|
||||
trajectory=result.get("trajectory"),
|
||||
trajectory_timesteps=result.get("trajectory_timesteps"),
|
||||
trajectory_decoded=result.get("trajectory_decoded"),
|
||||
video_path=result.get("video_path"),
|
||||
peak_memory_mb=result.get("peak_memory_mb"),
|
||||
state=result.get("state"),
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def to_legacy_dict(self) -> dict[str, Any]:
|
||||
result = {
|
||||
"prompts": self.prompt,
|
||||
"samples": self.samples,
|
||||
"frames": self.frames,
|
||||
"audio": self.audio,
|
||||
"size": self.size,
|
||||
"generation_time": self.generation_time,
|
||||
"logging_info": self.logging_info,
|
||||
"trajectory": self.trajectory,
|
||||
"trajectory_timesteps": self.trajectory_timesteps,
|
||||
"trajectory_decoded": self.trajectory_decoded,
|
||||
"video_path": self.video_path,
|
||||
"peak_memory_mb": self.peak_memory_mb,
|
||||
}
|
||||
if self.prompt_index is not None:
|
||||
result["prompt_index"] = self.prompt_index
|
||||
result["prompt"] = self.prompt
|
||||
if self.state is not None:
|
||||
result["state"] = self.state
|
||||
result.update(self.extra)
|
||||
return result
|
||||
|
||||
|
||||
__all__ = ["GenerationResult"]
|
||||
@@ -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,10 +84,40 @@ 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
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
|
||||
# GEN3C camera control
|
||||
trajectory_type: str | None = None
|
||||
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])
|
||||
|
||||
# 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
|
||||
@@ -86,26 +132,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:
|
||||
@@ -0,0 +1,296 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerConfig:
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8000
|
||||
output_dir: str = "outputs/"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParallelismConfig:
|
||||
tp_size: int = -1
|
||||
sp_size: int = -1
|
||||
hsdp_replicate_dim: int = 1
|
||||
hsdp_shard_dim: int = -1
|
||||
dist_timeout: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class OffloadConfig:
|
||||
dit: bool = True
|
||||
dit_layerwise: bool = True
|
||||
text_encoder: bool = True
|
||||
image_encoder: bool = True
|
||||
vae: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompileConfig:
|
||||
"""Typed ``torch.compile`` configuration.
|
||||
|
||||
``backend``/``fullgraph``/``mode``/``dynamic`` are the four most
|
||||
common ``torch.compile`` knobs. ``extras`` holds any remaining
|
||||
``torch.compile`` kwargs (e.g. ``options``, ``disable``).
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
text_encoder_enabled: bool | None = None
|
||||
"""Whether ``torch.compile`` is applied to the text encoder. ``None``
|
||||
keeps the runtime default. The public ``FastVideoArgs`` adapter does
|
||||
not yet consume this flag; reserved so the realtime runtime upstream
|
||||
(PR 7.6) has a typed home for its ``enable_torch_compile_text_encoder``
|
||||
kwarg without routing through ``pipeline.experimental``."""
|
||||
backend: str | None = None
|
||||
fullgraph: bool | None = None
|
||||
mode: str | None = None
|
||||
dynamic: bool | None = None
|
||||
extras: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class QuantizationConfig:
|
||||
text_encoder_quant: str | None = None
|
||||
transformer_quant: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class EngineConfig:
|
||||
num_gpus: int = 1
|
||||
execution_backend: Literal["mp", "ray"] = "mp"
|
||||
parallelism: ParallelismConfig = field(default_factory=ParallelismConfig)
|
||||
offload: OffloadConfig = field(default_factory=OffloadConfig)
|
||||
compile: CompileConfig = field(default_factory=CompileConfig)
|
||||
enable_stage_verification: bool = True
|
||||
use_fsdp_inference: bool = False
|
||||
disable_autocast: bool = False
|
||||
quantization: QuantizationConfig | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ComponentConfig:
|
||||
config_root: str | None = None
|
||||
pipeline_config_path: str | None = None
|
||||
text_encoder_weights: str | None = None
|
||||
transformer_weights: str | None = None
|
||||
transformer_2_weights: str | None = None
|
||||
vae_weights: str | None = None
|
||||
upsampler_weights: str | None = None
|
||||
lora_path: str | None = None
|
||||
override_pipeline_cls_name: str | None = None
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineSelection:
|
||||
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
|
||||
preset: str | None = None
|
||||
preset_version: int | None = None
|
||||
components: ComponentConfig = field(default_factory=ComponentConfig)
|
||||
vae_tiling: bool | None = None
|
||||
"""Tile-based VAE decode. ``None`` keeps the model's default."""
|
||||
preset_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
experimental: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GeneratorConfig:
|
||||
model_path: str
|
||||
revision: str | None = None
|
||||
trust_remote_code: bool = False
|
||||
engine: EngineConfig = field(default_factory=EngineConfig)
|
||||
pipeline: PipelineSelection = field(default_factory=PipelineSelection)
|
||||
|
||||
|
||||
@dataclass
|
||||
class InputConfig:
|
||||
prompt_path: str | None = None
|
||||
image_path: str | list[str] | None = None
|
||||
video_path: str | list[str] | None = None
|
||||
pil_image: Any | None = None
|
||||
pose: str | None = None
|
||||
mouse_cond: Any | None = None
|
||||
keyboard_cond: Any | None = None
|
||||
grid_sizes: Any | None = None
|
||||
c2ws_plucker_emb: Any | None = None
|
||||
refine_from: str | None = None
|
||||
stage1_video: Any | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingConfig:
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 1024
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
fps: int = 24
|
||||
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
|
||||
true_cfg_scale: float | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestRuntimeConfig:
|
||||
enable_teacache: bool = False
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputConfig:
|
||||
output_path: str = "outputs/"
|
||||
output_video_name: str | None = None
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
return_state: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContinuationState:
|
||||
kind: str
|
||||
payload: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PlannedStage:
|
||||
name: str
|
||||
kind: str
|
||||
source: str | None = None
|
||||
overrides: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationPlan:
|
||||
stages: list[PlannedStage]
|
||||
final_stage: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationRequest:
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str | None = None
|
||||
inputs: InputConfig = field(default_factory=InputConfig)
|
||||
sampling: SamplingConfig = field(default_factory=SamplingConfig)
|
||||
runtime: RequestRuntimeConfig = field(default_factory=RequestRuntimeConfig)
|
||||
output: OutputConfig = field(default_factory=OutputConfig)
|
||||
stage_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
state: ContinuationState | None = None
|
||||
plan: GenerationPlan | None = None
|
||||
extensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RunConfig:
|
||||
generator: GeneratorConfig
|
||||
request: GenerationRequest
|
||||
|
||||
|
||||
@dataclass
|
||||
class WarmupConfig:
|
||||
enabled: bool = True
|
||||
prompt: str = ("A cinematic drone shot over coastal cliffs at sunrise, "
|
||||
"golden light, gentle ocean waves, ultra detailed")
|
||||
timeout_seconds: int = 2400
|
||||
|
||||
|
||||
@dataclass
|
||||
class GpuPoolConfig:
|
||||
num_workers: int | None = None
|
||||
enable_audio_reencode: bool = True
|
||||
conditioning_num_frames: int = 9
|
||||
conditioning_end_offset: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptEnhancerConfig:
|
||||
enabled: bool = False
|
||||
provider: Literal["cerebras", "groq"] = "cerebras"
|
||||
model: str = "gpt-oss-120b"
|
||||
timeout_ms: int = 20000
|
||||
system_prompt_dir: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptSafetyConfig:
|
||||
enabled: bool = False
|
||||
classifier_path: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamingConfig:
|
||||
session_timeout_seconds: int = 300
|
||||
generation_segment_cap: int = 6
|
||||
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
|
||||
warmup: WarmupConfig = field(default_factory=WarmupConfig)
|
||||
pool: GpuPoolConfig = field(default_factory=GpuPoolConfig)
|
||||
prompt: PromptEnhancerConfig = field(default_factory=PromptEnhancerConfig)
|
||||
safety: PromptSafetyConfig = field(default_factory=PromptSafetyConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServeConfig:
|
||||
"""Typed serve config loaded from ``fastvideo serve --config``.
|
||||
|
||||
``default_request`` is a full :class:`GenerationRequest` — the same type
|
||||
clients POST to ``/v1/videos``. At request time the server merges it into
|
||||
the incoming body as the operator-pinned baseline.
|
||||
|
||||
Important nuance: only fields the operator **explicitly wrote** in the
|
||||
serve YAML/JSON count as defaults. Although the in-memory object is
|
||||
fully populated (schema defaults fill every unset field), the merge
|
||||
walks ``_fastvideo_explicit_paths`` — populated during parse — so
|
||||
unset fields are *not* forced onto requests. Per-request precedence:
|
||||
|
||||
body (client-explicit) > default_request (operator-explicit)
|
||||
> hardcoded fallback (e.g. ``fps=24``)
|
||||
|
||||
See :func:`fastvideo.api.compat.explicit_request_updates` for the
|
||||
projection and ``entrypoints/openai/video_api.py::_build_generation_kwargs``
|
||||
for the merge.
|
||||
"""
|
||||
generator: GeneratorConfig
|
||||
server: ServerConfig = field(default_factory=ServerConfig)
|
||||
default_request: GenerationRequest = field(default_factory=GenerationRequest)
|
||||
streaming: StreamingConfig | None = None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
"ComponentConfig",
|
||||
"ContinuationState",
|
||||
"EngineConfig",
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"GpuPoolConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"PromptEnhancerConfig",
|
||||
"PromptSafetyConfig",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
"StreamingConfig",
|
||||
"WarmupConfig",
|
||||
]
|
||||
@@ -0,0 +1,738 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Bidirectional Sparse Attention (BSA) backend for FastVideo.
|
||||
|
||||
Pure-PyTorch reference implementation from:
|
||||
"Bidirectional Sparse Attention for Faster Video Diffusion Training"
|
||||
(arXiv:2509.01085)
|
||||
|
||||
BSA sparsifies both queries (pruning redundant tokens per block) and
|
||||
key-value pairs (keeping only relevant KV blocks per query block).
|
||||
|
||||
This is a training-free inference backend: it works with any model
|
||||
trained with full attention by applying BSA sparsity at inference time.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.distributed import get_sp_group
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_varlen_func_impl, )
|
||||
FLASH_ATTN_AVAILABLE = True
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func as flash_attn_varlen_func_impl
|
||||
FLASH_ATTN_AVAILABLE = True
|
||||
except ImportError:
|
||||
FLASH_ATTN_AVAILABLE = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
BSA_TILE_SIZE = (4, 4, 4)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cached index helpers (same pattern as VSA)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_tile_partition_indices(
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Map raster-order tokens to tile-contiguous order."""
|
||||
T, H, W = dit_seq_shape
|
||||
ts, hs, ws = tile_size
|
||||
indices = torch.arange(T * H * W, device=device, dtype=torch.long).reshape(T, H, W)
|
||||
ls = []
|
||||
for t in range(math.ceil(T / ts)):
|
||||
for h in range(math.ceil(H / hs)):
|
||||
for w in range(math.ceil(W / ws)):
|
||||
ls.append(indices[
|
||||
t * ts:min(t * ts + ts, T),
|
||||
h * hs:min(h * hs + hs, H),
|
||||
w * ws:min(w * ws + ws, W),
|
||||
].flatten())
|
||||
return torch.cat(ls, dim=0)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_reverse_tile_partition_indices(
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Inverse mapping: tile-contiguous order back to raster order."""
|
||||
return torch.argsort(get_tile_partition_indices(dit_seq_shape, tile_size, device))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BSA core operations
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _prune_queries(
|
||||
q_blocks: torch.Tensor,
|
||||
keep_ratio: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""
|
||||
Prune redundant query tokens within each block.
|
||||
|
||||
Scores tokens by cosine similarity to the block center.
|
||||
Keeps the LEAST similar (most informative) tokens.
|
||||
|
||||
Args:
|
||||
q_blocks: [B, N_heads, N_blocks, block_size, D]
|
||||
keep_ratio: fraction of tokens to keep
|
||||
|
||||
Returns:
|
||||
sparse_q: [B, N_heads, N_blocks, keep_size, D]
|
||||
keep_indices: [B, N_heads, N_blocks, keep_size]
|
||||
keep_size: int
|
||||
"""
|
||||
B, H, N, S, D = q_blocks.shape
|
||||
keep_size = max(1, int(S * keep_ratio))
|
||||
|
||||
if keep_size >= S:
|
||||
idx = torch.arange(S, device=q_blocks.device)
|
||||
idx = idx.view(1, 1, 1, S).expand(B, H, N, S)
|
||||
return q_blocks, idx, S
|
||||
|
||||
center_idx = S // 2
|
||||
center = q_blocks[:, :, :, center_idx:center_idx + 1, :]
|
||||
|
||||
q_norm = F.normalize(q_blocks, dim=-1)
|
||||
c_norm = F.normalize(center, dim=-1)
|
||||
similarity = (q_norm * c_norm).sum(dim=-1) # [B, H, N, S]
|
||||
|
||||
# lowest similarity = most distinctive = keep
|
||||
_, indices = similarity.topk(keep_size, dim=-1, largest=False)
|
||||
indices, _ = indices.sort(dim=-1)
|
||||
|
||||
idx_expand = indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
|
||||
sparse_q = torch.gather(q_blocks, 3, idx_expand)
|
||||
|
||||
return sparse_q, indices, keep_size
|
||||
|
||||
|
||||
def _select_kv_blocks(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
cumulative_threshold: float,
|
||||
min_kv_blocks: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Dynamically select KV blocks for each query block.
|
||||
|
||||
Mean-pools to block level, computes block attention scores,
|
||||
admits blocks in descending order until cumulative mass
|
||||
exceeds threshold.
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
cumulative_threshold: e.g. 0.9
|
||||
min_kv_blocks: minimum blocks to keep
|
||||
|
||||
Returns:
|
||||
kv_mask: [B, H, N, N] boolean
|
||||
"""
|
||||
B, H, N, _, D = sparse_q.shape
|
||||
|
||||
q_repr = sparse_q.mean(dim=3)
|
||||
k_repr = k_blocks.mean(dim=3)
|
||||
|
||||
scores = torch.matmul(q_repr, k_repr.transpose(-1, -2)) / (D**0.5)
|
||||
block_attn = F.softmax(scores, dim=-1)
|
||||
|
||||
sorted_attn, sorted_idx = block_attn.sort(dim=-1, descending=True)
|
||||
cumsum = sorted_attn.cumsum(dim=-1)
|
||||
|
||||
keep_sorted = torch.ones_like(cumsum, dtype=torch.bool)
|
||||
keep_sorted[..., 1:] = cumsum[..., :-1] < cumulative_threshold
|
||||
|
||||
min_mask = torch.zeros_like(keep_sorted)
|
||||
min_mask[..., :min(min_kv_blocks, N)] = True
|
||||
keep_sorted = keep_sorted | min_mask
|
||||
|
||||
kv_mask = torch.zeros_like(block_attn, dtype=torch.bool)
|
||||
kv_mask.scatter_(-1, sorted_idx, keep_sorted)
|
||||
|
||||
return kv_mask
|
||||
|
||||
|
||||
def _compute_sparse_attention(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute attention for each query block against selected KV blocks.
|
||||
|
||||
Handles per-batch and per-head KV masks correctly.
|
||||
Uses flash_attn_varlen_func when available on GPU.
|
||||
Falls back to pure-PyTorch reference on CPU.
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
v_blocks: [B, H, N, Sk, D]
|
||||
kv_mask: [B, H, N, N] boolean (per-batch, per-head)
|
||||
|
||||
Returns:
|
||||
output: [B, H, N, Sq, D]
|
||||
"""
|
||||
if FLASH_ATTN_AVAILABLE and sparse_q.is_cuda:
|
||||
return _compute_sparse_attention_flash(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
else:
|
||||
return _compute_sparse_attention_reference(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
|
||||
|
||||
def _compute_sparse_attention_reference(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Pure-PyTorch fallback with per-batch, per-head mask support."""
|
||||
B, H, N, Sq, D = sparse_q.shape
|
||||
output = torch.zeros_like(sparse_q)
|
||||
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for qb in range(N):
|
||||
selected = kv_mask[b, h, qb] # [N] boolean
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
# [num_sel * Sk, D]
|
||||
sel_k = k_blocks[b, h, sel_idx].reshape(-1, D)
|
||||
sel_v = v_blocks[b, h, sel_idx].reshape(-1, D)
|
||||
|
||||
q = sparse_q[b, h, qb] # [Sq, D]
|
||||
scores = torch.matmul(q, sel_k.transpose(-1, -2)) / (D**0.5)
|
||||
weights = F.softmax(scores, dim=-1)
|
||||
output[b, h, qb] = torch.matmul(weights, sel_v)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _compute_sparse_attention_flash(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
FlashAttention implementation with per-batch, per-head mask support.
|
||||
|
||||
Strategy: check if all heads share the same mask. If so, use a single
|
||||
FlashAttention call per batch (fast path). If not, process each head
|
||||
separately (correct path).
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
v_blocks: [B, H, N, Sk, D]
|
||||
kv_mask: [B, H, N, N] boolean
|
||||
|
||||
Returns:
|
||||
output: [B, H, N, Sq, D]
|
||||
"""
|
||||
B, H, N, Sq, D = sparse_q.shape
|
||||
Sk = k_blocks.shape[3]
|
||||
device = sparse_q.device
|
||||
output = torch.zeros_like(sparse_q)
|
||||
|
||||
for b in range(B):
|
||||
# Check if all heads share the same mask for this batch element
|
||||
# Compare each head's mask to head 0's mask
|
||||
head0_mask = kv_mask[b, 0] # [N, N]
|
||||
all_heads_same = all(torch.equal(kv_mask[b, h], head0_mask) for h in range(1, H))
|
||||
|
||||
if all_heads_same:
|
||||
# Fast path: all heads share the same mask, single FA call
|
||||
_flash_attn_single_mask(
|
||||
sparse_q[b],
|
||||
k_blocks[b],
|
||||
v_blocks[b],
|
||||
head0_mask,
|
||||
output[b],
|
||||
H,
|
||||
N,
|
||||
Sq,
|
||||
Sk,
|
||||
D,
|
||||
device,
|
||||
)
|
||||
else:
|
||||
# Per-head path: process each head individually
|
||||
for h in range(H):
|
||||
head_mask = kv_mask[b, h] # [N, N]
|
||||
# Process single head: squeeze head dim, run FA, put back
|
||||
_flash_attn_single_head(
|
||||
sparse_q[b, h],
|
||||
k_blocks[b, h],
|
||||
v_blocks[b, h],
|
||||
head_mask,
|
||||
output,
|
||||
b,
|
||||
h,
|
||||
N,
|
||||
Sq,
|
||||
Sk,
|
||||
D,
|
||||
device,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _flash_attn_single_mask(
|
||||
sparse_q_b: torch.Tensor, # [H, N, Sq, D]
|
||||
k_blocks_b: torch.Tensor, # [H, N, Sk, D]
|
||||
v_blocks_b: torch.Tensor, # [H, N, Sk, D]
|
||||
mask: torch.Tensor, # [N, N] boolean
|
||||
output_b: torch.Tensor, # [H, N, Sq, D] (modified in-place)
|
||||
H: int,
|
||||
N: int,
|
||||
Sq: int,
|
||||
Sk: int,
|
||||
D: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Run FlashAttention for all heads sharing the same KV mask."""
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
cu_seqlens_q = [0]
|
||||
cu_seqlens_k = [0]
|
||||
active_blocks = []
|
||||
|
||||
for qb in range(N):
|
||||
selected = mask[qb] # [N] boolean
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
active_blocks.append(qb)
|
||||
num_kv_tokens = sel_idx.shape[0] * Sk
|
||||
|
||||
# [H, Sq, D] -> [Sq, H, D]
|
||||
q_block = sparse_q_b[:, qb].permute(1, 0, 2)
|
||||
q_list.append(q_block)
|
||||
|
||||
# [H, num_sel, Sk, D] -> [num_kv_tokens, H, D]
|
||||
sel_k = k_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
|
||||
sel_v = v_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
|
||||
k_list.append(sel_k)
|
||||
v_list.append(sel_v)
|
||||
|
||||
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
|
||||
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
|
||||
|
||||
if not q_list:
|
||||
return
|
||||
|
||||
flat_q = torch.cat(q_list, dim=0)
|
||||
flat_k = torch.cat(k_list, dim=0)
|
||||
flat_v = torch.cat(v_list, dim=0)
|
||||
|
||||
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
|
||||
|
||||
max_seqlen_q = Sq
|
||||
max_seqlen_k = int((cu_seqlens_k_t[1:] - cu_seqlens_k_t[:-1]).max().item())
|
||||
|
||||
orig_dtype = flat_q.dtype
|
||||
compute_dtype = orig_dtype
|
||||
if compute_dtype not in (torch.float16, torch.bfloat16):
|
||||
compute_dtype = torch.bfloat16
|
||||
flat_q = flat_q.to(compute_dtype)
|
||||
flat_k = flat_k.to(compute_dtype)
|
||||
flat_v = flat_v.to(compute_dtype)
|
||||
|
||||
flat_out = flash_attn_varlen_func_impl(
|
||||
flat_q,
|
||||
flat_k,
|
||||
flat_v,
|
||||
cu_seqlens_q_t,
|
||||
cu_seqlens_k_t,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
if compute_dtype != orig_dtype:
|
||||
flat_out = flat_out.to(orig_dtype)
|
||||
|
||||
idx = 0
|
||||
for qb in active_blocks:
|
||||
block_out = flat_out[idx:idx + Sq] # [Sq, H, D]
|
||||
output_b[:, qb] = block_out.permute(1, 0, 2) # [H, Sq, D]
|
||||
idx += Sq
|
||||
|
||||
|
||||
def _flash_attn_single_head(
|
||||
sparse_q_bh: torch.Tensor, # [N, Sq, D]
|
||||
k_blocks_bh: torch.Tensor, # [N, Sk, D]
|
||||
v_blocks_bh: torch.Tensor, # [N, Sk, D]
|
||||
mask: torch.Tensor, # [N, N] boolean
|
||||
output: torch.Tensor, # [B, H, N, Sq, D] (modified in-place)
|
||||
b: int,
|
||||
h: int,
|
||||
N: int,
|
||||
Sq: int,
|
||||
Sk: int,
|
||||
D: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Run FlashAttention for a single head with its own KV mask."""
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
cu_seqlens_q = [0]
|
||||
cu_seqlens_k = [0]
|
||||
active_blocks = []
|
||||
|
||||
for qb in range(N):
|
||||
selected = mask[qb]
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
active_blocks.append(qb)
|
||||
num_kv_tokens = sel_idx.shape[0] * Sk
|
||||
|
||||
# [Sq, D] -> [Sq, 1, D] (single head)
|
||||
q_block = sparse_q_bh[qb].unsqueeze(1)
|
||||
q_list.append(q_block)
|
||||
|
||||
# [num_sel, Sk, D] -> [num_kv_tokens, 1, D]
|
||||
sel_k = k_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
|
||||
sel_v = v_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
|
||||
k_list.append(sel_k)
|
||||
v_list.append(sel_v)
|
||||
|
||||
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
|
||||
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
|
||||
|
||||
if not q_list:
|
||||
return
|
||||
|
||||
flat_q = torch.cat(q_list, dim=0)
|
||||
flat_k = torch.cat(k_list, dim=0)
|
||||
flat_v = torch.cat(v_list, dim=0)
|
||||
|
||||
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
|
||||
|
||||
max_seqlen_q = Sq
|
||||
max_seqlen_k = int((cu_seqlens_k_t[1:] - cu_seqlens_k_t[:-1]).max().item())
|
||||
|
||||
orig_dtype = flat_q.dtype
|
||||
compute_dtype = orig_dtype
|
||||
if compute_dtype not in (torch.float16, torch.bfloat16):
|
||||
compute_dtype = torch.bfloat16
|
||||
flat_q = flat_q.to(compute_dtype)
|
||||
flat_k = flat_k.to(compute_dtype)
|
||||
flat_v = flat_v.to(compute_dtype)
|
||||
|
||||
flat_out = flash_attn_varlen_func_impl(
|
||||
flat_q,
|
||||
flat_k,
|
||||
flat_v,
|
||||
cu_seqlens_q_t,
|
||||
cu_seqlens_k_t,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
if compute_dtype != orig_dtype:
|
||||
flat_out = flat_out.to(orig_dtype)
|
||||
|
||||
idx = 0
|
||||
for qb in active_blocks:
|
||||
block_out = flat_out[idx:idx + Sq] # [Sq, 1, D]
|
||||
output[b, h, qb] = block_out.squeeze(1) # [Sq, D]
|
||||
idx += Sq
|
||||
|
||||
|
||||
def _reconstruct_pruned(
|
||||
sparse_output: torch.Tensor,
|
||||
keep_indices: torch.Tensor,
|
||||
block_size: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Scatter sparse output back to full block size.
|
||||
Pruned positions get nearest kept token's output.
|
||||
|
||||
Handles per-batch, per-head indices correctly.
|
||||
|
||||
Args:
|
||||
sparse_output: [B, H, N, keep_size, D]
|
||||
keep_indices: [B, H, N, keep_size]
|
||||
block_size: original tokens per block
|
||||
|
||||
Returns:
|
||||
full_output: [B, H, N, block_size, D]
|
||||
"""
|
||||
B, H, N, keep_size, D = sparse_output.shape
|
||||
device = sparse_output.device
|
||||
|
||||
if keep_size >= block_size:
|
||||
return sparse_output
|
||||
|
||||
full_output = torch.zeros(B, H, N, block_size, D, device=device, dtype=sparse_output.dtype)
|
||||
|
||||
# Scatter kept tokens
|
||||
idx_expand = keep_indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
|
||||
full_output.scatter_(3, idx_expand, sparse_output)
|
||||
|
||||
# Fill pruned positions with nearest kept token (vectorized)
|
||||
all_pos = torch.arange(block_size, device=device)
|
||||
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for n in range(N):
|
||||
kept = keep_indices[b, h, n] # [keep_size]
|
||||
|
||||
# Distance from every position to every kept position
|
||||
dists = (all_pos.view(-1, 1) - kept.view(1, -1)).abs()
|
||||
nearest_local_idx = dists.argmin(dim=1) # [block_size]
|
||||
|
||||
# Identify pruned positions
|
||||
is_pruned = torch.ones(block_size, dtype=torch.bool, device=device)
|
||||
is_pruned[kept] = False
|
||||
pruned_indices = is_pruned.nonzero(as_tuple=True)[0]
|
||||
|
||||
if pruned_indices.numel() > 0:
|
||||
src_indices = nearest_local_idx[pruned_indices]
|
||||
full_output[b, h, n, pruned_indices] = sparse_output[b, h, n, src_indices]
|
||||
|
||||
return full_output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FastVideo backend classes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BSAAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = False
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "BSA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["BSAAttentionImpl"]:
|
||||
return BSAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["BSAAttentionMetadata"]:
|
||||
return BSAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["BSAAttentionMetadataBuilder"]:
|
||||
return BSAAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class BSAAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
dit_seq_shape: tuple[int, int, int]
|
||||
total_seq_length: int
|
||||
num_blocks: int
|
||||
block_size: int
|
||||
tile_partition_indices: torch.LongTensor
|
||||
reverse_tile_partition_indices: torch.LongTensor
|
||||
# BSA-specific config
|
||||
query_keep_ratio: float
|
||||
kv_cumulative_threshold: float
|
||||
min_kv_blocks: int
|
||||
|
||||
|
||||
class BSAAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
bsa_query_keep_ratio: float = 0.5,
|
||||
bsa_kv_cumulative_threshold: float = 0.9,
|
||||
bsa_min_kv_blocks: int = 4,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> "BSAAttentionMetadata":
|
||||
# Ensure patching does not drop tokens silently.
|
||||
assert all(r % p == 0 for r, p in zip(raw_latent_shape, patch_size, strict=False)), (
|
||||
"raw_latent_shape must be divisible by patch_size for BSA", )
|
||||
|
||||
dit_seq_shape = (
|
||||
raw_latent_shape[0] // patch_size[0],
|
||||
raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2],
|
||||
)
|
||||
|
||||
total_seq_length = math.prod(dit_seq_shape)
|
||||
block_size = math.prod(BSA_TILE_SIZE)
|
||||
# Require exact tiling to avoid reshape failures later.
|
||||
assert all(d % t == 0 for d, t in zip(dit_seq_shape, BSA_TILE_SIZE, strict=False)), (
|
||||
"dit_seq_shape must be divisible by BSA_TILE_SIZE", )
|
||||
num_blocks = total_seq_length // block_size
|
||||
|
||||
tile_partition_indices = get_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
|
||||
reverse_tile_partition_indices = get_reverse_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
|
||||
|
||||
return BSAAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
dit_seq_shape=dit_seq_shape,
|
||||
total_seq_length=total_seq_length,
|
||||
num_blocks=num_blocks,
|
||||
block_size=block_size,
|
||||
tile_partition_indices=tile_partition_indices,
|
||||
reverse_tile_partition_indices=reverse_tile_partition_indices,
|
||||
query_keep_ratio=bsa_query_keep_ratio,
|
||||
kv_cumulative_threshold=bsa_kv_cumulative_threshold,
|
||||
min_kv_blocks=bsa_min_kv_blocks,
|
||||
)
|
||||
|
||||
|
||||
class BSAAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.prefix = prefix
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
if num_kv_heads is not None and num_kv_heads != num_heads:
|
||||
raise ValueError("BSA backend does not support grouped-query attention")
|
||||
if causal:
|
||||
raise ValueError("BSA backend is bidirectional; causal=True is unsupported")
|
||||
if softmax_scale is not None:
|
||||
expected_scale = 1.0 / math.sqrt(self.head_size)
|
||||
if not math.isclose(softmax_scale, expected_scale, rel_tol=1e-4, abs_tol=1e-5):
|
||||
raise ValueError("softmax_scale must be default (1/sqrt(d)) for BSA")
|
||||
try:
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
except (AssertionError, RuntimeError):
|
||||
self.sp_size = 1
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Reorder tokens from raster order to tile-contiguous order."""
|
||||
# qkv: [B, L, num_heads, D]
|
||||
return qkv[:, attn_metadata.tile_partition_indices]
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Reorder tokens from tile-contiguous order back to raster order."""
|
||||
return output[:, attn_metadata.reverse_tile_partition_indices]
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
BSA attention forward pass.
|
||||
|
||||
Input tensors are already in tile-contiguous order from preprocess_qkv.
|
||||
|
||||
Args:
|
||||
query: [B, L, num_heads, D] (tile-ordered)
|
||||
key: [B, L, num_heads, D] (tile-ordered)
|
||||
value: [B, L, num_heads, D] (tile-ordered)
|
||||
attn_metadata: BSA metadata
|
||||
|
||||
Returns:
|
||||
output: [B, L, num_heads, D] (tile-ordered)
|
||||
"""
|
||||
B, L, H, D = query.shape
|
||||
block_size = attn_metadata.block_size
|
||||
num_blocks = attn_metadata.num_blocks
|
||||
assert num_blocks * block_size == L, "Sequence length must match tiling"
|
||||
|
||||
# Reshape to [B, H, L, D] for attention computation
|
||||
q = query.transpose(1, 2).contiguous() # [B, H, L, D]
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Reshape into blocks: [B, H, num_blocks, block_size, D]
|
||||
q_blocks = q.view(B, H, num_blocks, block_size, D)
|
||||
k_blocks = k.view(B, H, num_blocks, block_size, D)
|
||||
v_blocks = v.view(B, H, num_blocks, block_size, D)
|
||||
|
||||
# --- Query sparsification ---
|
||||
sparse_q, keep_indices, keep_size = _prune_queries(q_blocks, attn_metadata.query_keep_ratio)
|
||||
|
||||
# --- KV block selection ---
|
||||
kv_mask = _select_kv_blocks(
|
||||
sparse_q,
|
||||
k_blocks,
|
||||
attn_metadata.kv_cumulative_threshold,
|
||||
attn_metadata.min_kv_blocks,
|
||||
)
|
||||
|
||||
# --- Sparse attention ---
|
||||
sparse_output = _compute_sparse_attention(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
|
||||
# --- Reconstruct pruned positions ---
|
||||
full_output = _reconstruct_pruned(sparse_output, keep_indices, block_size)
|
||||
|
||||
# Reshape back: [B, H, num_blocks, block_size, D] -> [B, H, L, D] -> [B, L, H, D]
|
||||
hidden_states = full_output.view(B, H, L, D).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,188 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_transformer_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CArchConfig(DiTArchConfig):
|
||||
"""Configuration for GEN3C architecture (VideoExtendGeneralDIT)."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Official GEN3C checkpoint key naming to FastVideo mapping.
|
||||
# The official checkpoint uses nn.Sequential patterns like attn.to_q.0 (Linear)
|
||||
# and attn.to_q.1 (RMSNorm), and layer1/layer2 for MLP.
|
||||
#
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$": r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding: net.t_embedder.1.linear_*.weight -> time_embed.t_embedder.linear_*.weight
|
||||
r"^net\.t_embedder\.0\.(.*)$": r"time_embed.time_proj.\1",
|
||||
r"^net\.t_embedder\.1\.linear_1\.(.*)$": r"time_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.t_embedder\.1\.linear_2\.(.*)$": r"time_embed.t_embedder.linear_2.\1",
|
||||
|
||||
# Augment sigma embedding (GEN3C-specific)
|
||||
r"^net\.augment_sigma_embedder\.0\.(.*)$": r"augment_sigma_embed.time_proj.\1",
|
||||
r"^net\.augment_sigma_embedder\.1\.linear_1\.(.*)$": r"augment_sigma_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.augment_sigma_embedder\.1\.linear_2\.(.*)$": r"augment_sigma_embed.t_embedder.linear_2.\1",
|
||||
|
||||
# Affine embedding norm: net.affline_norm.weight -> affine_norm.weight
|
||||
# Note: "affline" is a typo in the official GEN3C checkpoint (should be "affine")
|
||||
r"^net\.affline_norm\.(.*)$": r"affine_norm.\1",
|
||||
|
||||
# Extra positional embeddings (learnable per-axis)
|
||||
r"^net\.extra_pos_embedder\.pos_emb_t$": r"learnable_pos_embed.pos_emb_t",
|
||||
r"^net\.extra_pos_embedder\.pos_emb_h$": r"learnable_pos_embed.pos_emb_h",
|
||||
r"^net\.extra_pos_embedder\.pos_emb_w$": r"learnable_pos_embed.pos_emb_w",
|
||||
|
||||
# Transformer blocks: net.blocks.blockN -> transformer_blocks.N
|
||||
# Official uses: block.attn.to_q.0 (Linear), block.attn.to_q.1 (QK RMSNorm)
|
||||
#
|
||||
# Self-attention (block index 0)
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_out\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
# AdaLN modulation for self-attention
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.0\.adaLN_modulation\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.\2",
|
||||
|
||||
# Cross-attention (block index 1)
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_out\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
# AdaLN modulation for cross-attention
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.1\.adaLN_modulation\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.\2",
|
||||
|
||||
# MLP (block index 2): layer1 -> fc_in, layer2 -> fc_out
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer1\.(.*)$": r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer2\.(.*)$": r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
# AdaLN modulation for MLP
|
||||
r"^net\.blocks\.block(\d+)\.blocks\.2\.adaLN_modulation\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.\2",
|
||||
|
||||
# Final layer: net.final_layer.linear -> final_layer.proj_out
|
||||
r"^net\.final_layer\.linear\.(.*)$": r"final_layer.proj_out.\1",
|
||||
# Final layer AdaLN: net.final_layer.adaLN_modulation -> final_layer.adaln_modulation
|
||||
r"^net\.final_layer\.adaLN_modulation\.(.*)$": r"final_layer.adaln_modulation.\1",
|
||||
|
||||
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
|
||||
# - net.pos_embedder.* (rope position embeddings computed dynamically)
|
||||
# - net.accum_* keys (training metadata)
|
||||
# - logvar.* (training-only module, not used in inference)
|
||||
})
|
||||
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$": r"transformer_blocks.\1.mlp.\2",
|
||||
})
|
||||
|
||||
# GEN3C architecture parameters
|
||||
# Base VAE latent channels
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
|
||||
# Channels per 3D cache buffer: 16 (warped frame latent) + 16 (warped mask latent)
|
||||
CHANNELS_PER_BUFFER: int = 32
|
||||
|
||||
# Number of 3D cache buffers
|
||||
frame_buffer_max: int = 2
|
||||
|
||||
# Attention configuration (7B model: 32 heads x 128 dim = 4096 hidden)
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128 # 4096 / 32
|
||||
num_layers: int = 28
|
||||
mlp_ratio: float = 4.0
|
||||
|
||||
# Text encoder configuration
|
||||
text_embed_dim: int = 1024
|
||||
|
||||
# AdaLN-LoRA configuration
|
||||
adaln_lora_dim: int = 256
|
||||
use_adaln_lora: bool = True
|
||||
|
||||
# GEN3C-specific: augment sigma embedding for conditioning noise augmentation
|
||||
# Note: The official GEN3C-Cosmos-7B checkpoint was trained without this
|
||||
add_augment_sigma_embedding: bool = False
|
||||
|
||||
# Position embedding configuration
|
||||
max_size: tuple[int, int, int] = (128, 240, 240) # T, H, W
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0) # T, H, W scaling
|
||||
|
||||
# GEN3C uses learnable positional embeddings in addition to RoPE
|
||||
extra_pos_embed_type: str = "learnable"
|
||||
|
||||
# Padding mask handling
|
||||
concat_padding_mask: bool = True
|
||||
|
||||
# Cross-attention projection (not used in GEN3C 7B)
|
||||
use_crossattn_projection: bool = False
|
||||
|
||||
# RoPE FPS modulation
|
||||
rope_enable_fps_modulation: bool = True
|
||||
|
||||
# QK normalization
|
||||
qk_norm: str = "rms_norm"
|
||||
eps: float = 1e-6
|
||||
|
||||
# Affine embedding normalization
|
||||
affine_emb_norm: bool = True
|
||||
|
||||
# Block format (THWBD for GEN3C compatibility)
|
||||
block_x_format: str = "THWBD"
|
||||
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels
|
||||
|
||||
# Calculate total input channels for patch embedding:
|
||||
# - in_channels (16): VAE latent
|
||||
# - condition_video_input_mask (1): Binary mask for conditioning frames
|
||||
# - condition_video_pose (frame_buffer_max * 32): 3D cache buffers
|
||||
# - padding_mask (1 if concat_padding_mask): Padding mask
|
||||
self.buffer_channels = self.frame_buffer_max * self.CHANNELS_PER_BUFFER
|
||||
self.total_input_channels = (
|
||||
self.in_channels + # 16: VAE latent
|
||||
1 + # 1: condition_video_input_mask
|
||||
self.buffer_channels # 64: 3D cache buffers (2 * 32)
|
||||
)
|
||||
# padding_mask is added in build_patch_embed if concat_padding_mask=True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CVideoConfig(DiTConfig):
|
||||
"""Configuration for GEN3C video generation model."""
|
||||
arch_config: DiTArchConfig = field(default_factory=Gen3CArchConfig)
|
||||
prefix: str = "Gen3C"
|
||||
@@ -1,6 +1,7 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
|
||||
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
@@ -12,6 +13,7 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CVAEConfig(CosmosVAEConfig):
|
||||
"""
|
||||
GEN3C VAE config placeholder.
|
||||
|
||||
GEN3C uses tokenizer-backed VAE loading logic at runtime, but we keep a
|
||||
model-specific config class so pipeline/model configs stay model-scoped.
|
||||
"""
|
||||
@@ -5,7 +5,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig, WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.gen3c import Gen3CVideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.t5 import (T5LargeArchConfig, T5LargeConfig)
|
||||
from fastvideo.configs.models.vaes import Gen3CVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Gen3CT5LargeArchConfig(T5LargeArchConfig):
|
||||
"""T5 Large arch config that pads inputs to max_length.
|
||||
|
||||
GEN3C requires padded text encoder inputs, while the base
|
||||
T5 config no longer pads by default after the SP mask
|
||||
refactor [PR#1142](https://github.com/hao-ai-lab/FastVideo/pull/1142).
|
||||
"""
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Gen3CT5LargeConfig(T5LargeConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=_Gen3CT5LargeArchConfig)
|
||||
prefix: str = "t5"
|
||||
|
||||
|
||||
def t5_large_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""Postprocess T5 Large text encoder outputs for GEN3C pipeline.
|
||||
|
||||
Return raw last_hidden_state without truncation/padding.
|
||||
"""
|
||||
hidden_state = outputs.last_hidden_state
|
||||
|
||||
if hidden_state is None:
|
||||
raise ValueError("T5 Large outputs missing last_hidden_state")
|
||||
|
||||
nan_count = torch.isnan(hidden_state).sum()
|
||||
if nan_count > 0:
|
||||
hidden_state = hidden_state.masked_fill(torch.isnan(hidden_state), 0.0)
|
||||
|
||||
# Zero out embeddings beyond actual sequence length (vectorized)
|
||||
if outputs.attention_mask is not None:
|
||||
attention_mask = outputs.attention_mask
|
||||
lengths = attention_mask.sum(dim=1)
|
||||
max_len = hidden_state.shape[1]
|
||||
mask = torch.arange(max_len, device=hidden_state.device)[None, :] >= lengths[:, None]
|
||||
hidden_state[mask] = 0.0
|
||||
|
||||
return hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CConfig(PipelineConfig):
|
||||
"""Configuration for GEN3C Video Generation Pipeline.
|
||||
|
||||
GEN3C extends Cosmos with 3D cache for camera-controlled video generation.
|
||||
Key parameters:
|
||||
- frame_buffer_max: Number of 3D cache buffers (default: 2)
|
||||
- noise_aug_strength: Strength of noise augmentation per buffer
|
||||
- filter_points_threshold: Threshold for filtering unreliable depth points
|
||||
"""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Gen3CVideoConfig)
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=Gen3CVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (_Gen3CT5LargeConfig(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (t5_large_postprocess_text, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# GEN3C-specific conditioning parameters
|
||||
conditioning_strategy: str = "frame_replace"
|
||||
min_num_conditional_frames: int = 1
|
||||
max_num_conditional_frames: int = 2
|
||||
# Match official GEN3C/Cosmos inference defaults.
|
||||
sigma_conditional: float = 0.001
|
||||
sigma_data: float = 0.5
|
||||
state_ch: int = 16
|
||||
state_t: int = 16 # GEN3C uses 16 latent frames (121 pixel frames)
|
||||
text_encoder_class: str = "T5"
|
||||
|
||||
# Flow matching parameters
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: float = 1.0
|
||||
|
||||
# GEN3C 3D Cache parameters
|
||||
frame_buffer_max: int = 2
|
||||
noise_aug_strength: float = 0.0
|
||||
filter_points_threshold: float = 0.05
|
||||
|
||||
# Depth estimation settings
|
||||
use_moge_depth: bool = True
|
||||
moge_model_name: str = "Ruicheng/moge-vitl"
|
||||
offload_moge_after_depth: bool = True
|
||||
|
||||
# Camera trajectory settings (matching NVIDIA inference defaults)
|
||||
default_trajectory_type: str = "left"
|
||||
default_movement_distance: float = 0.3
|
||||
default_camera_rotation: str = "center_facing"
|
||||
|
||||
# Video generation settings
|
||||
# Match official GEN3C defaults (height=704, width=1280).
|
||||
video_resolution: tuple[int, int] = (704, 1280) # H, W
|
||||
num_frames: int = 121 # Default number of frames to generate
|
||||
|
||||
# Generation frame rate
|
||||
fps: int = 24
|
||||
|
||||
# Explicit CFG behavior policy:
|
||||
# - "legacy": CFG branch only when guidance_scale > 1.0
|
||||
# - "official_uncond_at_unity": also run uncond branch at guidance_scale == 1.0
|
||||
cfg_behavior: str = "legacy"
|
||||
default_negative_prompt: str = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality.")
|
||||
|
||||
# Autoregressive generation settings
|
||||
autoregressive_chunk_frames: int = 121 # Frames per chunk
|
||||
autoregressive_overlap_frames: int = 1 # Overlap between chunks
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
self._vae_latent_dim = 16
|
||||
|
||||
# Validate frame buffer configuration matches DiT
|
||||
if hasattr(self.dit_config, 'arch_config'):
|
||||
arch_config = self.dit_config.arch_config
|
||||
if (hasattr(arch_config, 'frame_buffer_max') and arch_config.frame_buffer_max != self.frame_buffer_max):
|
||||
raise ValueError(f"frame_buffer_max mismatch: pipeline config has {self.frame_buffer_max}, "
|
||||
f"DiT config has {arch_config.frame_buffer_max}")
|
||||
|
||||
allowed_cfg_behavior = {"legacy", "official_uncond_at_unity"}
|
||||
if self.cfg_behavior not in allowed_cfg_behavior:
|
||||
raise ValueError(f"cfg_behavior must be one of {sorted(allowed_cfg_behavior)}, got {self.cfg_behavior!r}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3CInferenceConfig(Gen3CConfig):
|
||||
"""Configuration for GEN3C inference with optimized defaults."""
|
||||
|
||||
# Use smaller batch sizes for inference
|
||||
batch_size: int = 1
|
||||
|
||||
# Enable gradient checkpointing for memory efficiency
|
||||
gradient_checkpointing: bool = False
|
||||
|
||||
# Inference-specific parameters
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 35
|
||||
|
||||
# Disable noise augmentation during inference
|
||||
noise_aug_strength: float = 0.0
|
||||
@@ -11,10 +11,34 @@ import torch
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5ArchConfig(T5ArchConfig):
|
||||
"""T5 arch that pads tokenizer output to ``max_length``.
|
||||
|
||||
LongCat's denoising stage concatenates positive and negative
|
||||
attention masks along the batch dimension for CFG, which requires
|
||||
uniform seq length. The shared :class:`T5ArchConfig` dropped the
|
||||
``"padding": "max_length"`` tokenizer kwarg so other DiTs could run
|
||||
with variable-length masks; LongCat still needs the uniform
|
||||
contract.
|
||||
"""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5Config(T5Config):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LongCatT5ArchConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields."""
|
||||
@@ -103,8 +127,9 @@ class LongCatT2V480PConfig(PipelineConfig):
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
# UMT5 uses T5-like config; postprocess pads to 512. LongCatT5Config
|
||||
# restores ``padding="max_length"`` for the CFG concat contract.
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (LongCatT5Config(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (longcat_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (umt5_postprocess_text, ))
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.hunyuangamecraft import (
|
||||
HunyuanGameCraftSamplingParam,
|
||||
HunyuanGameCraft65FrameSamplingParam,
|
||||
HunyuanGameCraft129FrameSamplingParam,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SamplingParam",
|
||||
"HunyuanGameCraftSamplingParam",
|
||||
"HunyuanGameCraft65FrameSamplingParam",
|
||||
"HunyuanGameCraft129FrameSamplingParam",
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos_Predict2_2B_Video2World_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 93
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = "The video captures a series of frames showing ugly scenes, static with no motion, motion blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of poor quality."
|
||||
num_inference_steps: int = 35
|
||||
@@ -1,23 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25SamplingParamBase(SamplingParam):
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 77
|
||||
fps: int = 24
|
||||
seed: int = 0
|
||||
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, "
|
||||
"low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, "
|
||||
"unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
|
||||
"Overall, the video is of poor quality.")
|
||||
num_inference_steps: int = 35
|
||||
@@ -1,21 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanSamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastHunyuanSamplingParam(HunyuanSamplingParam):
|
||||
num_inference_steps: int = 6
|
||||
@@ -1,55 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import numpy as np
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_480P_SamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 121
|
||||
height: int = 480
|
||||
width: int = 848
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 6.0
|
||||
sigmas: list[float] | None = field(default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
|
||||
|
||||
negative_prompt: str = ""
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.sigmas = list(np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_480P_StepDistilled_I2V_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
num_inference_steps: int = 12
|
||||
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_720P_Distilled_I2V_SamplingParam(Hunyuan15_720P_SamplingParam):
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_SR_1080P_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
|
||||
num_inference_steps: int = 12
|
||||
num_inference_steps_sr: int = 8
|
||||
|
||||
guidance_scale: float = 1.0
|
||||
@@ -1,92 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Sampling parameters for HunyuanGameCraft video generation.
|
||||
|
||||
GameCraft generates game-like videos with camera/action control.
|
||||
Default parameters are based on the official implementation.
|
||||
"""
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanGameCraftSamplingParam(SamplingParam):
|
||||
"""Sampling parameters for HunyuanGameCraft video generation.
|
||||
|
||||
Supports camera/action conditioning via:
|
||||
- camera_trajectory: Plücker coordinates for camera motion
|
||||
- action_list: List of actions (e.g., ["forward", "left", "right"])
|
||||
- action_speed_list: Speed multipliers for each action
|
||||
|
||||
Default resolution is 704x1280 (same as HunyuanVideo).
|
||||
Default frame count is 33 video frames -> 9 latent frames.
|
||||
"""
|
||||
|
||||
# Number of denoising steps
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Video dimensions
|
||||
# 33 video frames -> 9 latent frames (4x temporal compression)
|
||||
num_frames: int = 33
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
fps: int = 24
|
||||
|
||||
# Guidance scale - official GameCraft uses CFG with guidance_scale=6.0
|
||||
guidance_scale: float = 6.0
|
||||
|
||||
# Negative prompt for CFG (empty string = unconditional)
|
||||
negative_prompt: str = ""
|
||||
|
||||
# Camera/Action conditioning
|
||||
# Camera states as Plücker coordinates [B, T_video, 6, H, W]
|
||||
camera_states: Any | None = None
|
||||
|
||||
# Camera trajectory file/identifier (alternative to camera_states)
|
||||
camera_trajectory: str | None = None
|
||||
|
||||
# Action list for camera motion (e.g., ["forward", "left"])
|
||||
action_list: list[str] | None = None
|
||||
|
||||
# Speed multipliers for each action
|
||||
action_speed_list: list[float] | None = None
|
||||
|
||||
# History frame conditioning (for autoregressive generation)
|
||||
# Ground truth latents for conditioning [B, 16, T, H, W]
|
||||
gt_latents: Any | None = None
|
||||
|
||||
# Mask for conditioning (1=use gt, 0=generate) [B, 1, T, H, W]
|
||||
conditioning_mask: Any | None = None
|
||||
|
||||
# Number of conditioning frames (for autoregressive) - maps to num_cond_frames
|
||||
num_cond_frames: int = 0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# Validate action lists
|
||||
if (self.action_list is not None and self.action_speed_list is not None
|
||||
and len(self.action_list) != len(self.action_speed_list)):
|
||||
raise ValueError(f"action_list length ({len(self.action_list)}) must match "
|
||||
f"action_speed_list length ({len(self.action_speed_list)})")
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanGameCraft65FrameSamplingParam(HunyuanGameCraftSamplingParam):
|
||||
"""Sampling parameters for 65-frame GameCraft generation.
|
||||
|
||||
65 video frames -> 17 latent frames (with first frame as key frame).
|
||||
This is useful for longer video generation.
|
||||
"""
|
||||
num_frames: int = 65
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanGameCraft129FrameSamplingParam(HunyuanGameCraftSamplingParam):
|
||||
"""Sampling parameters for 129-frame GameCraft generation.
|
||||
|
||||
129 video frames -> 33 latent frames.
|
||||
This is the maximum supported by the official implementation.
|
||||
"""
|
||||
num_frames: int = 129
|
||||
@@ -1,25 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
import numpy as np
|
||||
|
||||
|
||||
@dataclass
|
||||
class HYWorld_SamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 125
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
fps: int = 24
|
||||
|
||||
# Camera trajectory: pose string (e.g., 'w-31' means generating [1 + 31] latents) or JSON file path
|
||||
pose: str = 'w-31'
|
||||
|
||||
guidance_scale: float = 6.0
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
sigmas: list[float] | None = field(default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
|
||||
|
||||
negative_prompt: str = ""
|
||||
@@ -1,20 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld_SamplingParam(Wan2_2_I2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 5.0 # high_noise
|
||||
guidance_scale_2: float = 5.0 # low_noise
|
||||
num_inference_steps: int = 70
|
||||
boundary_ratio: float | None = 0.947
|
||||
negative_prompt: str | None = ("画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
|
||||
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
|
||||
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
|
||||
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
|
||||
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
|
||||
"皮肤,肢体,面部特征,汽车,电线")
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
@@ -1,69 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2BaseSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 base one-stage T2V.
|
||||
|
||||
Values follow the official LTX-2 one-stage defaults.
|
||||
Multi-modal CFG params are read by ``LTX2DenoisingStage``.
|
||||
"""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 512
|
||||
width: int = 768
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 40
|
||||
guidance_scale: float = 3.0
|
||||
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
|
||||
negative_prompt: str = ("blurry, out of focus, overexposed, underexposed, low contrast, "
|
||||
"washed out colors, excessive noise, grainy texture, poor lighting, "
|
||||
"flickering, motion blur, distorted proportions, unnatural skin "
|
||||
"tones, deformed facial features, asymmetrical face, missing facial "
|
||||
"features, extra limbs, disfigured hands, wrong hand count, "
|
||||
"artifacts around text, inconsistent perspective, camera shake, "
|
||||
"incorrect depth of field, background too sharp, background clutter, "
|
||||
"distracting reflections, harsh shadows, inconsistent lighting "
|
||||
"direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
|
||||
"wrong gender, exaggerated expressions, wrong gaze direction, "
|
||||
"mismatched lip sync, silent or muted audio, distorted voice, "
|
||||
"robotic voice, echo, background noise, off-sync audio, incorrect "
|
||||
"dialogue, added dialogue, repetitive speech, jittery movement, "
|
||||
"awkward pauses, incorrect timing, unnatural transitions, "
|
||||
"inconsistent framing, tilted camera, flat lighting, inconsistent "
|
||||
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
|
||||
# Official LTX-2 multi-modal CFG defaults.
|
||||
ltx2_cfg_scale_video: float = 3.0
|
||||
ltx2_cfg_scale_audio: float = 7.0
|
||||
ltx2_modality_scale_video: float = 3.0
|
||||
ltx2_modality_scale_audio: float = 3.0
|
||||
ltx2_rescale_scale: float = 0.7
|
||||
# STG (Spatio-Temporal Guidance) defaults from official LTX-2.
|
||||
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])
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2DistilledSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 1024
|
||||
width: int = 1536
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 8
|
||||
guidance_scale: float = 1.0
|
||||
# No default negative_prompt for distilled models
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
# Backward compatibility alias.
|
||||
LTX2SamplingParam = LTX2DistilledSamplingParam
|
||||
@@ -1,25 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class SD35SamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 512
|
||||
width: int = 512
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 6.0
|
||||
@@ -1,73 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
TurboDiffusion sampling parameters.
|
||||
|
||||
TurboDiffusion uses RCM (recurrent Consistency Model) scheduler for
|
||||
1-4 step video generation with no classifier-free guidance.
|
||||
"""
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class TurboDiffusionT2V_1_3B_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for TurboDiffusion T2V 1.3B model.
|
||||
|
||||
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
|
||||
"""
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
# No negative prompt needed for TurboDiffusion (no CFG)
|
||||
negative_prompt: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TurboDiffusionT2V_14B_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for TurboDiffusion T2V 14B model.
|
||||
|
||||
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
|
||||
"""
|
||||
# Video parameters (720p for 14B)
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
# No negative prompt needed for TurboDiffusion (no CFG)
|
||||
negative_prompt: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TurboDiffusionI2V_A14B_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for TurboDiffusion I2V A14B model.
|
||||
|
||||
Uses 4-step RCM sampling with dual-model switching (high/low noise).
|
||||
"""
|
||||
# Video parameters (720p for A14B I2V)
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
# Note: boundary_ratio is set in the pipeline config (TurboDiffusionI2VConfig),
|
||||
# not here. This keeps sampling params and pipeline config separate.
|
||||
|
||||
# No negative prompt needed for TurboDiffusion (no CFG)
|
||||
negative_prompt: str | None = None
|
||||
@@ -1,154 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V_1_3B_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V_14B_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V_14B_480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# DMD parameters
|
||||
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
|
||||
num_inference_steps: int = 3
|
||||
num_frames: int = 61
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_Control_SamplingParam(SamplingParam):
|
||||
fps: int = 16
|
||||
num_frames: int = 49
|
||||
height: int = 832
|
||||
width: int = 480
|
||||
guidance_scale: float = 6.0
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.2 TI2V Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_2_Base_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 121
|
||||
fps: int = 24
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 4.0 # high_noise
|
||||
guidance_scale_2: float = 3.0 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 3.5 # high_noise
|
||||
guidance_scale_2: float = 3.5 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_Fun_A14B_Control_SamplingParam(Wan2_1_Fun_1_3B_Control_SamplingParam):
|
||||
num_frames: int = 81
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Causal Self-Forcing =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(Wan2_1_Fun_1_3B_InP_SamplingParam):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(Wan2_2_T2V_A14B_SamplingParam):
|
||||
num_inference_steps: int = 8
|
||||
num_frames: int = 81
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatrixGame2_SamplingParam(SamplingParam):
|
||||
height: int = 352
|
||||
width: int = 640
|
||||
num_frames: int = 57
|
||||
fps: int = 25
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 3
|
||||
negative_prompt: str | None = None
|
||||
@@ -7,7 +7,7 @@ Example usage:
|
||||
# launch a server and benchmark on it
|
||||
|
||||
# T2V or T2I or any other multimodal generation model
|
||||
fastvideo serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers --port 8000
|
||||
fastvideo serve --config serve.yaml
|
||||
|
||||
# benchmark it and make sure the port is the same as the server's port
|
||||
fastvideo bench --dataset vbench --num-prompts 20 --port 8000
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
"""``fastvideo eval`` CLI: list registered eval metrics and run them
|
||||
against a set of videos.
|
||||
|
||||
This is a thin wrapper around :mod:`fastvideo.eval`. Heavy lifting
|
||||
(metric loading, GPU handling, batching) lives in
|
||||
:func:`fastvideo.eval.create_evaluator`.
|
||||
|
||||
Examples::
|
||||
|
||||
fastvideo eval list
|
||||
fastvideo eval list --group vbench
|
||||
fastvideo eval run --videos path/to/videos/*.mp4 \\
|
||||
--metrics common.ssim --reference path/to/refs/
|
||||
fastvideo eval run --videos clip.mp4 --metrics vbench.aesthetic_quality \\
|
||||
--output scores.json
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class EvalSubcommand(CLISubcommand):
|
||||
"""The ``eval`` subcommand — entry point for the eval suite."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.name = "eval"
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
action = getattr(args, "eval_action", None)
|
||||
if action == "list":
|
||||
_cmd_list(args)
|
||||
elif action == "run":
|
||||
_cmd_run(args)
|
||||
else:
|
||||
# Re-print help if no action was given.
|
||||
self._parser.print_help() # type: ignore[attr-defined]
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
action = getattr(args, "eval_action", None)
|
||||
if action == "run" and not args.videos:
|
||||
raise SystemExit("`fastvideo eval run` requires --videos")
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
eval_parser = subparsers.add_parser(
|
||||
"eval",
|
||||
help="Run video-gen evaluation metrics",
|
||||
usage="fastvideo eval {list,run} [...]",
|
||||
)
|
||||
sub = eval_parser.add_subparsers(dest="eval_action", required=False)
|
||||
|
||||
# `eval list`
|
||||
list_p = sub.add_parser("list", help="List registered metrics")
|
||||
list_p.add_argument("--group", type=str, default=None, help="Filter to a metric group (e.g. 'vbench').")
|
||||
|
||||
# `eval run`
|
||||
run_p = sub.add_parser("run", help="Evaluate videos against one or more metrics")
|
||||
run_p.add_argument("--videos",
|
||||
type=str,
|
||||
nargs="+",
|
||||
required=False,
|
||||
help="Path, glob, or directory of generated videos.")
|
||||
run_p.add_argument("--reference",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path / glob / dir of reference videos (for paired metrics).")
|
||||
run_p.add_argument("--metrics", type=str, default="all", help="Comma-separated metric names, or 'all'.")
|
||||
run_p.add_argument("--device", type=str, default="cuda", help="Torch device (e.g. 'cuda', 'cuda:0', 'cpu').")
|
||||
run_p.add_argument("--text-prompt",
|
||||
type=str,
|
||||
nargs="*",
|
||||
default=None,
|
||||
help="Prompt(s) for text-conditioned metrics. One per video.")
|
||||
run_p.add_argument("--fps", type=float, default=None, help="Frame-rate annotation passed to fps-aware metrics.")
|
||||
run_p.add_argument("--output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Write results as JSON to this path (default: stdout).")
|
||||
|
||||
# Stash the parser so cmd() can re-print help on no-action.
|
||||
self._parser = eval_parser # type: ignore[attr-defined]
|
||||
return cast(FlexibleArgumentParser, eval_parser)
|
||||
|
||||
|
||||
def _cmd_list(args: argparse.Namespace) -> None:
|
||||
from fastvideo.eval import list_metrics
|
||||
names = list_metrics()
|
||||
if args.group:
|
||||
prefix = args.group.rstrip(".") + "."
|
||||
names = [n for n in names if n == args.group or n.startswith(prefix)]
|
||||
if not names:
|
||||
print(f"(no metrics matched group {args.group!r})")
|
||||
return
|
||||
for name in names:
|
||||
print(name)
|
||||
|
||||
|
||||
def _cmd_run(args: argparse.Namespace) -> None:
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
video_paths = _expand_paths(args.videos)
|
||||
if not video_paths:
|
||||
raise SystemExit(f"No videos matched: {args.videos}")
|
||||
ref_paths = _expand_paths([args.reference]) if args.reference else None
|
||||
|
||||
metrics_arg: list[str] | str = ("all" if args.metrics == "all" else
|
||||
[m.strip() for m in args.metrics.split(",") if m.strip()])
|
||||
|
||||
evaluator = create_evaluator(metrics=metrics_arg, device=args.device)
|
||||
|
||||
all_results: list[dict] = []
|
||||
for i, vp in enumerate(video_paths):
|
||||
logger.info("Evaluating %s (%d/%d)", vp, i + 1, len(video_paths))
|
||||
kwargs: dict = {"video": load_video(vp)}
|
||||
if ref_paths is not None:
|
||||
ref = ref_paths[i] if i < len(ref_paths) else ref_paths[0]
|
||||
kwargs["reference"] = load_video(ref)
|
||||
if args.text_prompt is not None:
|
||||
prompt = (args.text_prompt[i] if i < len(args.text_prompt) else args.text_prompt[0])
|
||||
kwargs["text_prompt"] = [prompt]
|
||||
if args.fps is not None:
|
||||
kwargs["fps"] = args.fps
|
||||
|
||||
results = evaluator.evaluate(**kwargs)
|
||||
all_results.append({
|
||||
"video": str(vp),
|
||||
"scores": _serialize_results(results),
|
||||
})
|
||||
|
||||
payload = json.dumps(all_results, indent=2, default=_jsonable)
|
||||
if args.output:
|
||||
Path(args.output).write_text(payload)
|
||||
logger.info("Wrote results to %s", args.output)
|
||||
else:
|
||||
print(payload)
|
||||
|
||||
|
||||
def _expand_paths(patterns: list[str]) -> list[str]:
|
||||
out: list[str] = []
|
||||
for pat in patterns:
|
||||
p = Path(pat)
|
||||
if p.is_dir():
|
||||
for ext in (".mp4", ".avi", ".mov", ".mkv", ".gif"):
|
||||
out.extend(sorted(str(f) for f in p.iterdir() if f.suffix.lower() == ext))
|
||||
elif any(c in pat for c in "*?["):
|
||||
out.extend(sorted(glob.glob(pat)))
|
||||
else:
|
||||
out.append(pat)
|
||||
# de-dup, preserve order
|
||||
seen: set[str] = set()
|
||||
deduped: list[str] = []
|
||||
for x in out:
|
||||
if x not in seen:
|
||||
seen.add(x)
|
||||
deduped.append(x)
|
||||
return deduped
|
||||
|
||||
|
||||
def _serialize_results(results) -> dict | list:
|
||||
"""Turn evaluator output (dict or list-of-dicts) into JSON-friendly form."""
|
||||
if isinstance(results, list):
|
||||
return [_serialize_results(r) for r in results]
|
||||
if isinstance(results, dict):
|
||||
return {k: _serialize_metric_result(v) for k, v in results.items()}
|
||||
return _serialize_metric_result(results)
|
||||
|
||||
|
||||
def _serialize_metric_result(mr) -> dict:
|
||||
return {
|
||||
"name": getattr(mr, "name", None),
|
||||
"score": getattr(mr, "score", None),
|
||||
"details": getattr(mr, "details", None),
|
||||
}
|
||||
|
||||
|
||||
def _jsonable(obj):
|
||||
"""``json.dumps(default=...)`` coercer for metric outputs.
|
||||
|
||||
Metrics frequently land numpy scalars / arrays, torch tensors, and
|
||||
pathlib paths inside ``MetricResult.details`` (e.g. ``optical_flow``
|
||||
populates ``per_frame_metrics`` with numpy floats). The stdlib JSON
|
||||
encoder rejects all of those by default — this callback walks the
|
||||
leaves and coerces them to native Python types.
|
||||
"""
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
if isinstance(obj, np.integer):
|
||||
return int(obj)
|
||||
if isinstance(obj, np.floating):
|
||||
return float(obj)
|
||||
if isinstance(obj, np.bool_):
|
||||
return bool(obj)
|
||||
if isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
if isinstance(obj, torch.Tensor):
|
||||
return obj.detach().cpu().tolist()
|
||||
if isinstance(obj, Path):
|
||||
return str(obj)
|
||||
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
return [EvalSubcommand()]
|
||||
@@ -2,19 +2,17 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.utils import RaiseNotImplementedAction
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.entrypoints.cli.inference_config import build_generate_run_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_VALIDATED_RUN_CONFIG_ATTR = "_fastvideo_validated_run_config"
|
||||
|
||||
|
||||
class GenerateSubcommand(CLISubcommand):
|
||||
@@ -23,89 +21,47 @@ class GenerateSubcommand(CLISubcommand):
|
||||
def __init__(self) -> None:
|
||||
self.name = "generate"
|
||||
super().__init__()
|
||||
self.init_arg_names = self._get_init_arg_names()
|
||||
self.generation_arg_names = self._get_generation_arg_names()
|
||||
|
||||
def _get_init_arg_names(self) -> list[str]:
|
||||
"""Get names of arguments for VideoGenerator initialization"""
|
||||
return ["num_gpus", "tp_size", "sp_size", "model_path"]
|
||||
|
||||
def _get_generation_arg_names(self) -> list[str]:
|
||||
"""Get names of arguments for generate_video method"""
|
||||
return [field.name for field in dataclasses.fields(SamplingParam)]
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = ['subparser', 'config', 'dispatch_function']
|
||||
run_config = getattr(args, _VALIDATED_RUN_CONFIG_ATTR, None)
|
||||
if run_config is None:
|
||||
run_config = build_generate_run_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
)
|
||||
logger.info("CLI generate config: %s", run_config)
|
||||
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (k not in excluded_args and v is not None and hasattr(args, '_provided') and k in args._provided):
|
||||
provided_args[k] = v
|
||||
|
||||
if 'model_path' in vars(args) and args.model_path is not None:
|
||||
provided_args['model_path'] = args.model_path
|
||||
|
||||
if 'prompt' in vars(args) and args.prompt is not None:
|
||||
provided_args['prompt'] = args.prompt
|
||||
|
||||
merged_args = {**provided_args}
|
||||
|
||||
logger.info('CLI Args: %s', merged_args)
|
||||
|
||||
if 'model_path' not in merged_args or not merged_args['model_path']:
|
||||
raise ValueError("model_path must be provided either in config file or via --model-path")
|
||||
|
||||
# Check if either prompt or prompt_txt is provided
|
||||
has_prompt = 'prompt' in merged_args and merged_args['prompt']
|
||||
has_prompt_txt = 'prompt_txt' in merged_args and merged_args['prompt_txt']
|
||||
|
||||
if not (has_prompt or has_prompt_txt):
|
||||
raise ValueError("Either prompt or prompt_txt must be provided")
|
||||
|
||||
if has_prompt and has_prompt_txt:
|
||||
raise ValueError("Cannot provide both 'prompt' and 'prompt_txt'. Use only one of them.")
|
||||
|
||||
init_args = {k: v for k, v in merged_args.items() if k not in self.generation_arg_names}
|
||||
generation_args = {k: v for k, v in merged_args.items() if k in self.generation_arg_names}
|
||||
generation_args.setdefault("return_frames", False)
|
||||
|
||||
model_path = init_args.pop('model_path')
|
||||
prompt = generation_args.pop('prompt', None)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_path=model_path, **init_args)
|
||||
|
||||
# Call generate_video - it handles both single and batch modes
|
||||
generator.generate_video(prompt=prompt, **generation_args)
|
||||
generator = VideoGenerator.from_config(run_config.generator)
|
||||
generator.generate(run_config.request)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
"""Validate the arguments for this command"""
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
if args.config and not os.path.exists(args.config):
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo generate requires --config PATH; use a nested "
|
||||
"run config plus optional dotted overrides")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Config file not found: {args.config}")
|
||||
setattr(
|
||||
args,
|
||||
_VALIDATED_RUN_CONFIG_ATTR,
|
||||
build_generate_run_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
),
|
||||
)
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
generate_parser = subparsers.add_parser(
|
||||
"generate",
|
||||
help="Run inference on a model",
|
||||
usage="fastvideo generate (--model-path MODEL_PATH_OR_ID --prompt PROMPT) | --config CONFIG_FILE [OPTIONS]")
|
||||
usage="fastvideo generate --config RUN_CONFIG [--dotted.override VALUE]")
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default='',
|
||||
required=False,
|
||||
help="Read CLI options from a config JSON or YAML file. If provided, --model-path and --prompt are optional."
|
||||
)
|
||||
|
||||
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
|
||||
generate_parser = SamplingParam.add_cli_args(generate_parser)
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--text-encoder-configs",
|
||||
action=RaiseNotImplementedAction,
|
||||
help="JSON array of text encoder configurations (NOT YET IMPLEMENTED)",
|
||||
help="Path to a nested run config JSON or YAML file. Required.",
|
||||
)
|
||||
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.parser import load_raw_config, parse_config
|
||||
from fastvideo.api.schema import RunConfig, ServeConfig
|
||||
|
||||
_GENERATE_OVERRIDE_PREFIXES = ("generator.", "request.")
|
||||
_SERVE_OVERRIDE_PREFIXES = (
|
||||
"generator.",
|
||||
"server.",
|
||||
"default_request.",
|
||||
)
|
||||
|
||||
|
||||
def build_generate_run_config(
|
||||
args: argparse.Namespace,
|
||||
overrides: list[str] | None = None,
|
||||
) -> RunConfig:
|
||||
raw = _load_nested_config(getattr(args, "config", None))
|
||||
raw.setdefault("request", {})
|
||||
raw = _apply_dotted_overrides(
|
||||
raw,
|
||||
overrides,
|
||||
allowed_prefixes=_GENERATE_OVERRIDE_PREFIXES,
|
||||
)
|
||||
_ensure_generate_cli_defaults(raw)
|
||||
config = parse_config(RunConfig, raw)
|
||||
_validate_num_gpus(config.generator.engine.num_gpus)
|
||||
_validate_generate_prompt_sources(config)
|
||||
return config
|
||||
|
||||
|
||||
def build_serve_config(
|
||||
args: argparse.Namespace,
|
||||
overrides: list[str] | None = None,
|
||||
) -> ServeConfig:
|
||||
raw = _load_nested_config(getattr(args, "config", None))
|
||||
raw.setdefault("server", {})
|
||||
raw.setdefault("default_request", {})
|
||||
raw = _apply_dotted_overrides(
|
||||
raw,
|
||||
overrides,
|
||||
allowed_prefixes=_SERVE_OVERRIDE_PREFIXES,
|
||||
)
|
||||
config = parse_config(ServeConfig, raw)
|
||||
_validate_num_gpus(config.generator.engine.num_gpus)
|
||||
return config
|
||||
|
||||
|
||||
def _load_nested_config(path: str | None) -> dict[str, Any]:
|
||||
if not path:
|
||||
raise ValueError("Inference CLI requires --config PATH; use a nested config file "
|
||||
"plus optional dotted overrides")
|
||||
|
||||
raw = load_raw_config(path)
|
||||
if not isinstance(raw.get("generator"), Mapping):
|
||||
raise ValueError("Inference config must use the nested schema with a top-level "
|
||||
"'generator' mapping")
|
||||
return deepcopy(dict(raw))
|
||||
|
||||
|
||||
def _apply_dotted_overrides(
|
||||
raw: Mapping[str, Any],
|
||||
overrides: list[str] | None,
|
||||
*,
|
||||
allowed_prefixes: tuple[str, ...],
|
||||
) -> dict[str, Any]:
|
||||
if not overrides:
|
||||
return deepcopy(dict(raw))
|
||||
|
||||
parsed = parse_cli_overrides(overrides)
|
||||
for key in parsed:
|
||||
if "." not in key:
|
||||
raise ValueError("CLI overrides must use dotted config paths like "
|
||||
"--request.sampling.seed 42")
|
||||
if not key.startswith(allowed_prefixes):
|
||||
allowed = ", ".join(allowed_prefixes)
|
||||
raise ValueError(f"Unsupported override path {key!r}. Allowed prefixes: {allowed}")
|
||||
return apply_overrides(raw, parsed)
|
||||
|
||||
|
||||
def _ensure_generate_cli_defaults(raw: dict[str, Any]) -> None:
|
||||
request = raw.setdefault("request", {})
|
||||
output = request.setdefault("output", {})
|
||||
output.setdefault("return_frames", False)
|
||||
|
||||
|
||||
def _validate_generate_prompt_sources(config: RunConfig) -> None:
|
||||
has_prompt = config.request.prompt is not None
|
||||
has_prompt_path = config.request.inputs.prompt_path is not None
|
||||
if not (has_prompt or has_prompt_path):
|
||||
raise ValueError("Either request.prompt or request.inputs.prompt_path must be provided")
|
||||
if has_prompt and has_prompt_path:
|
||||
raise ValueError("Cannot provide both request.prompt and request.inputs.prompt_path")
|
||||
|
||||
|
||||
def _validate_num_gpus(num_gpus: int) -> None:
|
||||
if num_gpus <= 0:
|
||||
raise ValueError(f"generator.engine.num_gpus must be > 0; got {num_gpus}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_generate_run_config",
|
||||
"build_serve_config",
|
||||
]
|
||||
@@ -1,11 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
from fastvideo.entrypoints.cli.serve import cmd_init as serve_cmd_init
|
||||
from fastvideo.entrypoints.cli.bench import cmd_init as bench_cmd_init
|
||||
from fastvideo.entrypoints.cli.eval import cmd_init as eval_cmd_init
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
@@ -14,6 +14,7 @@ def cmd_init() -> list[CLISubcommand]:
|
||||
commands.extend(generate_cmd_init())
|
||||
commands.extend(serve_cmd_init())
|
||||
commands.extend(bench_cmd_init())
|
||||
commands.extend(eval_cmd_init())
|
||||
return commands
|
||||
|
||||
|
||||
@@ -27,14 +28,17 @@ def main() -> None:
|
||||
for cmd in cmd_init():
|
||||
cmd.subparser_init(subparsers).set_defaults(dispatch_function=cmd.cmd)
|
||||
cmds[cmd.name] = cmd
|
||||
args = parser.parse_args()
|
||||
|
||||
args, unknown = parser.parse_known_args()
|
||||
if unknown and args.subparser not in {"generate", "serve"}:
|
||||
parser.error(f"unrecognized arguments: {' '.join(unknown)}")
|
||||
args._unknown = unknown
|
||||
if args.subparser in cmds:
|
||||
cmds[args.subparser].validate(args)
|
||||
|
||||
if hasattr(args, "dispatch_function"):
|
||||
args.dispatch_function(args)
|
||||
else:
|
||||
parser.print_help()
|
||||
return
|
||||
|
||||
parser.print_help()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -2,14 +2,17 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.api.compat import generator_config_to_fastvideo_args
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.entrypoints.cli.inference_config import build_serve_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_VALIDATED_SERVE_CONFIG_ATTR = "_fastvideo_validated_serve_config"
|
||||
|
||||
|
||||
class ServeSubcommand(CLISubcommand):
|
||||
@@ -20,94 +23,69 @@ class ServeSubcommand(CLISubcommand):
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = {
|
||||
"subparser",
|
||||
"config",
|
||||
"dispatch_function",
|
||||
"host",
|
||||
"port",
|
||||
"output_dir",
|
||||
}
|
||||
serve_config = getattr(args, _VALIDATED_SERVE_CONFIG_ATTR, None)
|
||||
if serve_config is None:
|
||||
serve_config = build_serve_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
)
|
||||
|
||||
provided: set[str] = getattr(args, '_provided', set())
|
||||
cli_kwargs = {}
|
||||
for k, v in vars(args).items():
|
||||
if k in excluded_args:
|
||||
continue
|
||||
if k == '_provided':
|
||||
continue
|
||||
if k in provided and v is not None:
|
||||
cli_kwargs[k] = v
|
||||
logger.info("CLI serve config: %s", serve_config)
|
||||
|
||||
if 'model_path' not in cli_kwargs and args.model_path is not None:
|
||||
cli_kwargs['model_path'] = args.model_path
|
||||
|
||||
if not cli_kwargs.get('model_path'):
|
||||
raise ValueError("model_path must be provided via --model-path")
|
||||
# A `streaming:` block selects the WebSocket/Dynamo runtime;
|
||||
# its deps stay out of REST-only deployments via lazy import.
|
||||
if serve_config.streaming is not None:
|
||||
from fastvideo.entrypoints.streaming.server import (
|
||||
run_server as run_streaming_server, )
|
||||
run_streaming_server(serve_config)
|
||||
return
|
||||
|
||||
from fastvideo.entrypoints.openai.api_server import (
|
||||
DEFAULT_HOST,
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
DEFAULT_PORT,
|
||||
run_server,
|
||||
run_server, )
|
||||
|
||||
logger.info(
|
||||
"Server will listen on %s:%d",
|
||||
serve_config.server.host,
|
||||
serve_config.server.port,
|
||||
)
|
||||
|
||||
host = getattr(args, "host", DEFAULT_HOST)
|
||||
port = getattr(args, "port", DEFAULT_PORT)
|
||||
output_dir = getattr(args, "output_dir", DEFAULT_OUTPUT_DIR)
|
||||
|
||||
logger.info("CLI serve args: %s", cli_kwargs)
|
||||
logger.info("Server will listen on %s:%d", host, port)
|
||||
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(**cli_kwargs)
|
||||
run_server(fastvideo_args, host=host, port=port, output_dir=output_dir)
|
||||
fastvideo_args = generator_config_to_fastvideo_args(serve_config.generator)
|
||||
run_server(
|
||||
fastvideo_args,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
output_dir=serve_config.server.output_dir,
|
||||
default_request=serve_config.default_request,
|
||||
)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
from fastvideo.entrypoints.openai.api_server import (
|
||||
DEFAULT_HOST,
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
DEFAULT_PORT,
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo serve requires --config PATH; use a nested "
|
||||
"serve config plus optional dotted overrides")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Config file not found: {args.config}")
|
||||
setattr(
|
||||
args,
|
||||
_VALIDATED_SERVE_CONFIG_ATTR,
|
||||
build_serve_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
),
|
||||
)
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
serve_parser = subparsers.add_parser(
|
||||
"serve",
|
||||
help="Start an OpenAI-compatible HTTP server",
|
||||
usage=("fastvideo serve --model-path MODEL_PATH_OR_ID "
|
||||
"[--host HOST] [--port PORT] [OPTIONS]"),
|
||||
)
|
||||
|
||||
serve_parser.add_argument(
|
||||
"--host",
|
||||
type=str,
|
||||
default=DEFAULT_HOST,
|
||||
help=f"Host to bind the server to (default: {DEFAULT_HOST})",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--port",
|
||||
type=int,
|
||||
default=DEFAULT_PORT,
|
||||
help=f"Port to listen on (default: {DEFAULT_PORT})",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--output-dir",
|
||||
type=str,
|
||||
default=DEFAULT_OUTPUT_DIR,
|
||||
help=("Directory for generated outputs "
|
||||
f"(default: {DEFAULT_OUTPUT_DIR})"),
|
||||
usage="fastvideo serve --config SERVE_CONFIG [--dotted.override VALUE]",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="Read CLI options from a config JSON or YAML file.",
|
||||
help="Path to a nested config JSON or YAML file. Required.",
|
||||
)
|
||||
|
||||
serve_parser = FastVideoArgs.add_cli_args(serve_parser)
|
||||
return cast(FlexibleArgumentParser, serve_parser)
|
||||
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user