Compare commits

..
Author SHA1 Message Date
SolitaryThinker 49a13e6359 [skill] Add add-model and review-add-model-pr skills
Two skills split out of the will/stable-audio PR (#1260) so the model
work and the agent infra evolve on independent review cycles.

`add-model/SKILL.md` (~1600 lines) — the canonical procedure for
porting a new diffusion model into FastVideo, distilled from the Wan,
LTX-2, Hunyuan, Cosmos, Stable Audio, and MagiHuman ports. Covers the
17-row Files-table, 16-step procedure, FastVideo layer/attention
selection rules, parity-test pattern, weight conversion recipe, and
the I2V variant pattern.

`add-model/REVIEW.md` (~1900 lines) — 36 documented failure modes
discovered while porting / re-reviewing those families. Each item is
a self-contained "what we did wrong, what to do instead" entry with a
suggested action on the skill text. Items 1-22 from the Wan / MagiHuman
porting passes; items 23-36 from the Stable Audio port (audio as
first-class workload, hard ban on `from diffusers import` for model
classes at runtime, user-story example docstrings, single-class
kwargs-driven variants, examples accept paths not user-written decode
glue, etc.).

`add-model/add_model_split.md` — plan for breaking the monolithic
1600-line SKILL.md into ~12 satellite docs (deferred until after
the next porting pass validates which sections are most-referenced).

`review-add-model-pr/SKILL.md` (~480 lines, NEW) — reviewer-facing
dual of `add-model`. Walks a PR reviewer through the canonical
surface (per-component checklist for DiT / VAE / encoder /
PipelineConfig / pipeline class / stages / presets+registry / tests /
example), indexes the 36 REVIEW failure modes by where they show up
in a diff (the "Pitfall map"), and produces a structured verdict
(block / nit / follow-up) with a copy-pasteable comment skeleton.
Key design choice: the reviewer's job is NOT to redesign the port,
just to verify the porter followed the canonical procedure honestly
and exercised the parity gate — REVIEW item 16 ("skip-on-missing
parity is a silent no-op trap") gets disproportionate weight because
it's the single most expensive failure mode in past ports.

Both skills registered in `.agents/skills/index.jsonl` as
`status: draft, trust: low` per existing convention.
2026-04-26 17:14:49 -07:00
Mook 7b872cc41e [Perf] Skip bool-mask round-trip in block-sparse VSA attention (#1243) 2026-04-26 15:14:37 -07:00
alexzms 37418946c8 [docs]: clarify real_score_guidance_scale CFG parameterization (#1256) 2026-04-26 16:38:00 +08:00
William Lin 95fd29e0cb [feat] Streaming WebSocket server skeleton (single generator + fMP4) (#1251) 2026-04-26 00:33:49 -07:00
Junda Suandmergify[bot] e17cd2633c [bugfix]: normalize uint8 pil_image in I2V VAE encoding (#1249)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-24 09:16:01 +00:00
William Lin e0dc5f2b0c [feat] Add typed LTX-2 continuation state and streaming session store (#1250) 2026-04-24 01:28:07 -07:00
William Lin 70ee5d230c [feat] [6/n] Improve API: LTX-2 public preset + asset wiring + gpu_pool translation (#1239) 2026-04-23 11:36:45 -07:00
William Lin 24ced500f5 [test] add LTX-2 distilled T2V SSIM regression test (#1240) 2026-04-21 12:03:38 -07:00
William Lin 4ddcdf541f [feat] [5.5/n] Improve API: streaming server config surface + serve dispatch (#1238) 2026-04-17 15:36:21 -07:00
William Lin 0e3529869c [feat] [5/n] Improve API: wire ServeConfig.default_request into OpenAI serving (#1237) 2026-04-17 13:26:18 -07:00
William Lin e1e0d91c00 [misc] small cleanup for API handling (#1235) 2026-04-16 16:21:21 -07:00
William Lin 145a3f166b [feat] [4/n] Improve API: refactor sampling param and merge with presets (#1234) 2026-04-16 14:10:02 -07:00
William Lin 88a5a933ab [feat] [3/n] Improve API: extend support to cli (#1226) 2026-04-14 15:20:47 -07:00
William Lin c591d6d2a6 [feat] [2/n] Improve API: add initial support in video_generator (#1220) 2026-04-06 10:33:54 -07:00
Kun Linandmergify[bot] 65dff806a8 [bugfix]Fixing Lora distillation training distributed checkpointing bug (#1192)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-06 02:20:26 +00:00
KUAN-HAO HUANGandmergify[bot] b85f0f4c2a [perf]: Eliminate CPU-GPU synchronization bottlenecks in training pipeline (#1217)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-06 02:03:46 +00:00
William Lin 76c62d7a00 [feat] [1/n] API improvements: add intial files for new fastvideo public API (#1218) 2026-04-05 18:13:19 -07:00
f6e65ff668 [Feature] Add BSA (Bidirectional Sparse Attention) inference backend (#1174)
Co-authored-by: Satyam Srivastava <satyam53@Mac.lan1>
Co-authored-by: Satyam Srivastava <satyam53@Satyams-MacBook-Air.local>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-05 05:00:33 +00:00
mergify[bot] c220aa8000 [ci](mergify): upgrade configuration to current format (#1216)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-04 23:09:17 +00:00
Jinzhe PanandDarren Sadr 4713fc17ed [feat] Job Runner UI (#1189)
Co-authored-by: Darren Sadr <darrensadr@gmail.com>
2026-04-02 16:07:24 -07:00
vishruthb 5789955bbe [feat] add gen3c (cosmos-7b) model and pipeline support (#1059) 2026-04-01 11:42:02 +00:00
Jinzhe Pan 2ad84a3b78 [ci] Use update instead of rebase for auto branch sync (#1215) 2026-04-01 19:16:59 +08:00
Jinzhe Pan 12d699cd78 [ci] Add direct test retry with check overwrite and aggregate status refresh (#1214) 2026-04-01 17:21:28 +08:00
Jinzhe Pan 34f14ded21 [ci] Use pull_request_target for Full Suite trigger (#1213) 2026-04-01 03:01:07 +08:00
Jinzhe Pan 71d1ab411f [ci] Fix jq crash when Buildkite build env is null (#1212) 2026-04-01 02:35:01 +08:00
Jinzhe Pan 805e487773 [ci] Ignore legacy reference videos when checking for HF download (#1211) 2026-04-01 02:12:09 +08:00
Jinzhe Pan 8803b4547e [ci] Add retry for flaky tests and fix stale SSIM references (#1210) 2026-04-01 01:11:49 +08:00
Jinzhe Pan 3b3806b3f6 [ci] Fix /merge to directly trigger Full Suite + simplify rebase conditions (#1209) 2026-03-31 23:17:09 +08:00
Jinzhe Pan 38d962e89d [ci] Remove Mergify ready-label race condition (#1208) 2026-03-31 20:59:13 +08:00
Jinzhe Pan 3966a365d0 [ci] Add statuses:write permission for /test pre-commit (#1207) 2026-03-31 20:33:18 +08:00
Jinzhe Pan d73fd14af0 [ci] Post pre-commit status to PR commit SHA (#1206) 2026-03-31 20:21:21 +08:00
302 changed files with 37193 additions and 2013 deletions
+96
View File
@@ -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"
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+230
View File
@@ -0,0 +1,230 @@
# add-model skill — split plan
The current `SKILL.md` is **1605 lines** in a single file. This document
proposes splitting it into ~12 satellite docs with a much shorter
top-level index, following the idiomatic Anthropic-skills pattern of
"short procedural index + targeted satellites."
## Why split
1. **Cold-load cost.** The whole 1605-line file is loaded into the
model's context every time the skill fires. Most porting sessions
only need a fraction of that — e.g. a VAE-only contribution doesn't
need the I2V-variant section or the weight-conversion recipe. A
shorter `SKILL.md` index + lazy-loaded satellites keeps the live
context lean.
2. **Reviewability.** A single 1605-line markdown file is hard to
diff. Splits keep changes scoped — adding the "audio workload"
section (REVIEW item 25) becomes a new file in `how-to/` rather
than a 100-line insert in the middle of an existing megafile.
3. **Findability.** Section anchors in a single file are hard to
navigate; per-topic files surface in `ls .agents/skills/add-model/`
and `grep -r` returns clean per-file matches.
4. **Component-only contributions** (REVIEW item 23) become a
first-class workflow document instead of an awkward "but skip
half the steps" exception inside the main file.
## Current section inventory
| Section | Approx lines | Stays in SKILL.md? | Target file |
|---|---|---|---|
| Purpose | 8 | Yes | SKILL.md |
| When to use / not to use | 25 | Yes | SKILL.md |
| Prerequisites | 40 | Yes | SKILL.md |
| Inputs (table) | 15 | Yes | SKILL.md |
| FastVideo's single architecture | 40 | No | how-to/architecture.md |
| Files you will create or touch (table) | 25 | Yes (link to how-to) | SKILL.md + how-to/files_table.md |
| Steps (1–16) | 440 | **Index only** in SKILL.md (3-line per step + link) | how-to/steps_<phase>.md ×4 |
| Standard stages — subclass targets | 20 | No | how-to/architecture.md |
| FastVideo layers and attention | 110 | No | how-to/layers_and_attention.md |
| Parity test pattern | 165 | No | how-to/parity_testing.md |
| Parallel component porting | 100 | No | how-to/parity_testing.md |
| Weight conversion | 155 | No | how-to/weight_conversion.md |
| `register_configs` cheatsheet | 25 | No | how-to/registry_cheatsheet.md |
| Adding an I2V variant | 270 | No | how-to/i2v_variant.md |
| Distributed support | 35 | No | how-to/distributed_support.md |
| Common pitfalls | 50 | No | reference/pitfalls.md |
| Outputs | 15 | Yes | SKILL.md |
| Example prompt snippet | 35 | No | reference/example_prompt.md |
| References | 50 | No | reference/links.md |
| Changelog | 20 | No | reference/changelog.md |
## Target tree
```
.agents/skills/add-model/
├── SKILL.md ~250 lines — procedural index
├── REVIEW.md unchanged
├── add_model_split.md this doc
├── how-to/
│ ├── architecture.md ~70 lines — single-architecture diagram + standard stages
│ ├── files_table.md ~50 lines — annotated 17-row Files table
│ ├── steps_1_setup.md ~80 lines — steps 1–5 (gather, study, reuse, convert, clone)
│ ├── steps_2_components.md ~120 lines — step 6 + parallel component porting
│ ├── steps_3_pipeline.md ~100 lines — steps 7–11 (config, stages, pipeline class, presets, registry)
│ ├── steps_4_validation.md ~120 lines — steps 12–16 (smoke, pipeline parity, SSIM, cleanup, ask)
│ ├── parity_testing.md ~250 lines — conventions + component template + pipeline-level gate + subagent prompt
│ ├── weight_conversion.md ~200 lines — decision tree + recipe + reference scripts + gotchas
│ ├── layers_and_attention.md ~150 lines — linear/attention/primitive selection rules
│ ├── i2v_variant.md ~270 lines — full I2V add section verbatim (cleanly extractable)
│ ├── distributed_support.md ~50 lines — SP/TP/VAE-tiling rules
│ ├── registry_cheatsheet.md ~30 lines — register_configs fields
│ ├── component_only_contributions.md ~120 lines — NEW (REVIEW #23): VAE-only / encoder-only PR shape
│ └── audio_workload.md ~150 lines — NEW (REVIEW #25): audio output, T2A workload, no-SSIM metrics
├── reference/
│ ├── pitfalls.md ~80 lines — the 14-item pitfalls list
│ ├── example_prompt.md ~40 lines — the example user prompt snippet
│ ├── links.md ~50 lines — file/repo references
│ └── changelog.md ~30 lines — change history table
└── seed-ssim-references/ unchanged
```
Total: ~2200 lines across 17 files (vs current 1605 lines × 1 file).
The line growth is intentional — each file gets a short "Purpose +
Status + Prerequisites" header so it can be loaded without context
from the others.
## SKILL.md target shape (~250 lines)
```markdown
# Add a Model to FastVideo
## Purpose
[8 lines, unchanged]
## When to use / When not to use
[25 lines, unchanged]
## Prerequisites — gather inputs (blocking)
[40 lines, unchanged]
## Inputs
[15-line table, unchanged]
## Files you will create or touch
The full table lives in `how-to/files_table.md`. Quick sketch:
- Component files (model + config + __init__ export): rows 1–6.
- Pipeline files (config + class + stages + presets): rows 7–10.
- Registry: row 11.
- Tests (smoke + pipeline parity + per-component parity + SSIM): rows 12–15.
- Conversion script (only if not Diffusers format): row 16.
- Example: row 17.
For component-only contributions (just a VAE / encoder), see
`how-to/component_only_contributions.md` — you can skip rows 7–13 + 15
+ 17.
## Steps (procedural index)
1. **Gather inputs (blocking).** See `how-to/steps_1_setup.md`.
2. **Study the reference implementation.** See `how-to/steps_1_setup.md`.
3. **Decide what to reuse.** See `how-to/steps_1_setup.md`.
4. **Convert weights to Diffusers format.** Only if not Diffusers-format.
See `how-to/weight_conversion.md`.
5. **Clone the official repo for parity testing.** See `how-to/steps_1_setup.md`.
6. **Port components in parallel via subagents.** See
`how-to/steps_2_components.md` + `how-to/parity_testing.md`.
7. **Create the PipelineConfig.** See `how-to/steps_3_pipeline.md`.
8. **Build or pick the stages.** See `how-to/architecture.md` (standard
stages catalog) + `how-to/steps_3_pipeline.md`.
9. **Write the pipeline class.** See `how-to/steps_3_pipeline.md`.
10. **Define presets.** See `how-to/steps_3_pipeline.md`.
11. **Register in `fastvideo/registry.py`.** See `how-to/registry_cheatsheet.md`.
12. **Smoke-test the pipeline.** See `how-to/steps_4_validation.md`.
13. **Full-pipeline parity + example (gated).** See
`how-to/parity_testing.md` (the pipeline-level section is the
handoff gate).
14. **Add SSIM regression.** See `how-to/steps_4_validation.md`. For
audio, see `how-to/audio_workload.md` (no SSIM analog).
15. **Clean up the cloned reference repo.** See `how-to/steps_4_validation.md`.
16. **Ask about tests + perf data.** See `how-to/steps_4_validation.md`.
## Pre-handoff checklist
[New section addressing REVIEW item 16 — bans skip-only parity at handoff]
- [ ] `pytest tests/local_tests/<bucket>/test_<family>_*parity*.py -v` produces non-skip PASS for each non-reused component.
- [ ] `pytest tests/local_tests/pipelines/test_<family>_pipeline_parity.py -v` produces non-skip PASS.
- [ ] `python examples/inference/basic/basic_<family>.py` writes a non-corrupt mp4 (or .wav for audio).
- [ ] Conversion has actually been run (the parity tests skip if not).
## Outputs
[15 lines, unchanged]
## Common pitfalls
See `reference/pitfalls.md` for the full 14-item list. Most-cited:
- #1: `EntryClass` missing → pipeline silently invisible.
- #11: raw `nn.Linear` in DiT/VAE hot paths → use `ReplicatedLinear`.
- #16: skip-only parity → see pre-handoff checklist above.
## See also
- `reference/example_prompt.md` — example user prompt for invoking this skill.
- `reference/links.md` — file/repo references.
- `reference/changelog.md` — change history.
```
## Migration steps
1. **Create the new files** with content extracted from current
`SKILL.md` (no semantic edits — just relocation + cross-link edits).
2. **Add per-file headers** of the form:
```markdown
# <topic>
**Part of:** add-model skill
**When to read:** <one-liner — e.g. "during step 6 / parallel
component porting">
**Prerequisites:** <links to other files needed first>
---
<content>
```
3. **Rewrite SKILL.md** as the index. Each step gets the new 3-line
format: title + 2-line summary + link to `how-to/<file>.md`.
4. **Add the pre-handoff checklist** (addresses REVIEW item 16) — the
single load-bearing addition this split enables.
5. **Add `how-to/component_only_contributions.md`** (addresses REVIEW
item 23).
6. **Add `how-to/audio_workload.md`** (addresses REVIEW items 25 + 28).
7. **Update REVIEW.md** to mark items 16, 23, 25, 28 as resolved by
the split.
8. **Re-run a sample skill invocation** (e.g. on a hypothetical new
port) to validate the split — check that the model only loads
the index + 2-3 satellite files for a typical port.
## Risks / open questions
- **Breaks existing prompts.** Anyone with an in-flight skill
invocation may have memorized section anchors that move. Mitigate
by leaving anchor stubs in SKILL.md for one cycle (forwarding
comments).
- **Cross-link maintenance.** Each rename / move requires updating
cross-refs. Mitigate by keeping the file tree shallow (one
`how-to/` and one `reference/` directory only).
- **Discoverability of new files.** A porter who only reads
`SKILL.md` may not realize `audio_workload.md` exists. Mitigate by
having the index's step-3 line for "audio variant of step X" link
explicitly to the audio doc, and by keeping the "See also" section
visible.
- **What counts as "idiomatic"?** Anthropic skills tend toward
~150-300 line single-file or ~3-5 file splits. 17 files is on the
large side. We might collapse further if some satellites are
always-loaded-together (e.g. merge `architecture.md` +
`registry_cheatsheet.md` if they're never read independently).
Decide post-prototype.
## Next step
Implement the split as a separate PR (don't bundle with the
will/stable-audio first-class VAE work or the will/magi MagiHuman
port). Sequence:
1. Land the split as-is (mechanical relocation, zero semantic
change). Verify that running the skill against a known port
produces equivalent guidance.
2. Then layer in the REVIEW-item edits (16, 23, 25, 28) as content
changes inside the new structure.
+2
View File
@@ -5,3 +5,5 @@
{"name": "evaluate-video-quality", "description": "Evaluate generated video quality using available metrics (SSIM, loss trajectory, caption consistency)", "path": "evaluate-video-quality/SKILL.md", "status": "draft", "trust": "low"}
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
{"name": "review-add-model-pr", "description": "Review a PR that adds a new model / pipeline (or major variant) to FastVideo. Walks the canonical surface, indexes the 36 documented add-model failure modes by where they show up in a diff, and produces a structured verdict (block / nit / follow-up).", "path": "review-add-model-pr/SKILL.md", "status": "draft", "trust": "low"}
+482
View File
@@ -0,0 +1,482 @@
---
name: review-add-model-pr
description: Use when reviewing a PR that adds a new model / pipeline (or a major variant like I2V/V2V/DMD) to FastVideo under `fastvideo/pipelines/basic/<family>/`. Walks a reviewer through the canonical surface the porter should have touched, the parity bar, and the 36 documented failure modes from prior ports. Returns a structured review verdict (block / nit / follow-up).
---
# Review an add-model PR
## Purpose
Catch the failure modes that have actually shipped in past add-model PRs
(documented in `.agents/skills/add-model/REVIEW.md`) **before merge**,
without re-litigating design choices that the `add-model` skill has
already settled. The reviewer's job is *not* to redesign the port — it
is to verify the porter did the things the skill prescribes, exercised
the parity gate honestly, and surfaced the right knobs to end users.
This skill assumes the PR claims to add a model. For PRs adding only
sampling-preset tweaks, a single pipeline kwarg, or a new test against
an existing pipeline, this skill is overkill — comment scoped to those
specific changes instead.
## When to use
- A new pipeline directory under `fastvideo/pipelines/basic/<family>/`.
- A new variant pipeline (e.g. `<family>_i2v_pipeline.py`,
`<family>_dmd_pipeline.py`) sibling-added to an existing family.
- A "first-class component" PR that adds a new DiT, VAE, or encoder
port without (yet) wiring a pipeline (REVIEW item 23).
- Re-review of a port that previously skipped the post-parity hot-path
pass (REVIEW item 33) or shipped with placeholder diffusers imports
(REVIEW item 30).
## When not to use
- PR only edits sampling defaults in an existing
`<family>/presets.py` → review the values against the model card
inline, no skill needed.
- PR only adds a new SSIM reference → use the
`seed-ssim-references` skill instead.
- PR refactors shared infra (`fastvideo/layers/`, `fastvideo/attention/`,
`fastvideo/pipelines/stages/` base classes) without touching a
family directory → that's not an add-model PR.
## Required reading before starting the review
1. **The PR description.** What family is being added, what variants
(T2V / I2V / V2V / DMD / T2A / …), what the porter claims is parity-
verified, and which of the four prereqs (official repo URL, HF
weights path, HF token, target `model_family`) they collected.
2. **`.agents/skills/add-model/SKILL.md` Files-table** (rows 1–17) — the
canonical surface a model port touches. You'll cross-reference this
against the diff in step 2 below.
3. **`.agents/skills/add-model/REVIEW.md` summary table** — the 36
failure modes. The "Pitfall map" section below indexes them by where
they typically show up in a diff.
## Inputs
| Parameter | Required | Description |
|-----------|----------|-------------|
| `pr` | Yes | GitHub PR number, URL, or local branch name. |
| `model_family` | Recommended | The snake_case family slug (e.g. `stable_audio`, `magi_human`). Lets you grep-scope the review. |
| `official_repo` | Recommended | URL of the upstream reference, so you can sanity-check the parity test references it correctly. |
| `parity_run_log` | Optional | If the porter attached a parity run log, read it before reading code. |
## Steps
### 1. Verify the PR scope is shaped like an add-model PR
Run `gh pr view <pr> --json files | jq -r '.files[].path'` (or
`git diff --name-only origin/main...HEAD`). Confirm the diff touches
**at least 6 of the 17 Files-table rows** in `add-model/SKILL.md`.
Typical shape:
| If you see... | Expect... | If missing... |
|---|---|---|
| `fastvideo/models/dits/<family>.py` | `fastvideo/configs/models/dits/<family>.py` + export in `__init__.py` | Block: DiT config + export missing (Files-table rows 2 + 3). |
| `fastvideo/pipelines/basic/<family>/<family>_pipeline.py` | `presets.py`, `registry.py` edit, smoke test, parity test, example | Block on whichever is missing — the porter will be told the same by the skill. |
| `<family>_i2v_pipeline.py` (sibling-added) | A new preset for the I2V workload + the registry entry pointing at the I2V HF repo | Block: I2V variants get their own preset + workload tag (REVIEW item 8). |
| Standalone DiT/VAE/encoder under `fastvideo/models/<bucket>/` with no pipeline | A "first-class component" PR (REVIEW item 23) | Confirm with the author that the consuming pipeline PR is in flight or planned. Don't block. |
| Edits under `fastvideo/configs/pipelines/base.py`, `fastvideo/api/sampling_param.py`, `fastvideo/pipelines/pipeline_batch_info.py` | New pipeline-call kwargs the family needs | Sanity-check the field's default doesn't change behavior for other families. Often a footgun (e.g. `0 or default` Python truthiness). |
### 2. Walk the diff in dependency order, not file order
Components have a strict dependency order — review them in this order
so you can spot mismatches as they cascade:
1. **DiT** (`fastvideo/models/dits/<family>.py` + config) — foundational.
2. **VAE** — only one usually exists; if not, this is the second-biggest
review surface.
3. **Encoder(s) / conditioner** — text encoder is usually shared; new
conditioners (e.g. `MultiConditioner`-style) need careful look.
4. **`PipelineConfig`** subclass.
5. **Pipeline class** — wires modules + stages.
6. **Stages** — most should be standard; mod-specific subclasses
warrant careful review.
7. **Presets + registry** — usually mechanical; check workload type
selection against REVIEW item 28.
8. **Tests** — both component parity and pipeline parity.
9. **Example script** — last, but the bar is high (REVIEW items 31, 35).
For each component, run the Per-component checklist (next section).
### 3. Apply the Per-component checklist
For each new file in `fastvideo/models/<bucket>/<family>.py`:
#### DiT review (`fastvideo/models/dits/<family>.py`)
- [ ] **No raw `nn.Linear`** in QKV/MLP/proj/embedder paths. All
projections should be `ReplicatedLinear` from
`fastvideo.layers.linear`. Exceptions are the legacy MoE-packed
case (REVIEW item 11) — must be flagged in a comment with the
weight-layout justification.
- [ ] **No raw SDPA / `flash_attn_func`** calls. Self-attention should
be `DistributedAttention` (or `LocalAttention` for cross-attention).
See `add-model/SKILL.md` "Attention layers" table.
- [ ] **No raw `nn.LayerNorm` in modulation paths.** Should be
`FP32LayerNorm` / `RMSNorm` / `LayerNormScaleShift` from
`fastvideo.layers.layernorm`.
- [ ] **No `from diffusers import` or `from transformers import
<ModelClass>`** anywhere outside of test files (REVIEW item 30, the
hard ban). Tokenizers are the *only* allowed `from transformers`
runtime import in production code.
- [ ] **`from_official_state_dict()` (or equivalent)** is provided so
the consuming pipeline can load the published checkpoint without
going through `diffusers.from_pretrained`.
- [ ] **Param-mapping deviations from upstream are localized** — e.g.
if the model uses `gamma`/`beta` and FastVideo's `FP32LayerNorm`
uses `weight`/`bias`, the remap happens *in the loader*, not by
renaming the layers.
- [ ] **Partial / non-standard rotary** (e.g. halves-swap vs
interleaved-pair) is kept local with a one-line WHY comment, not a
copy of FastVideo's `_apply_rotary_emb` with edits.
#### VAE review (`fastvideo/models/vaes/<family or arch>.py`)
- [ ] **File name matches the convention** (REVIEW item 29): name
after the *arch* if the VAE is shared across families
(`oobleck.py`, `autoencoder_kl.py`); name after the *family* if it's
specific (`wanvae.py`).
- [ ] **`from_pretrained()` (or `from_official_state_dict`)** loads
weights from the published HF path *without* a Diffusers
intermediary (REVIEW item 30 again).
- [ ] **Per-channel `latents_mean`/`latents_std` are reshaped with
explicit `.view(1, -1, 1, 1, 1)`** when applied (REVIEW item 22) —
silent broadcasting along the wrong dim is a stealth bug.
- [ ] **Normalization-convention mismatches with the wrapper**
(REVIEW item 20) — if upstream `decode()` does
`z = z*std + mean` *internally* but the FastVideo wrapper expects
pre-denormalized input, the parity test must compensate or the
decode parity will look broken.
- [ ] **Pipeline-glue wrapper** (e.g. `fastvideo/models/vaes/<family>_audio.py`
for lazy-load semantics) exists if the pipeline needs lazy-load /
hide-from-named_parameters semantics (REVIEW item 26).
#### Conditioner / encoder review
- [ ] If the conditioner mixes text + numeric (e.g. duration), it
produces the **DiT-ready (cross_attn_cond, cross_attn_mask,
global_embed) triple** in a single helper, not via inline cat-ing
in the conditioning stage.
- [ ] **T5 / Llama / SigLIP TP wiring** uses the encoder bucket's TP
primitives (`QKVParallelLinear`, `MergedColumnParallelLinear`,
`RowParallelLinear`) — not `ReplicatedLinear`.
- [ ] If the conditioner intentionally hides its T5 from
`named_parameters()` (so the SA-style checkpoint loader doesn't try
to match upstream-absent T5 keys), that exclusion is documented in
a comment.
#### `PipelineConfig` review (`fastvideo/configs/pipelines/<family>.py`
or `fastvideo/pipelines/basic/<family>/pipeline_configs.py`)
- [ ] **Subclasses `PipelineConfig`** from
`fastvideo.configs.pipelines.base`.
- [ ] **`vae_config`, `dit_config`, `text_encoder_configs`** are
defaulted via `field(default_factory=...)` (mutable-default-arg
Python rule).
- [ ] **The text-encoder slot** matches what the pipeline actually
uses. If the pipeline owns its own conditioner (no FastVideo-loaded
text encoder), the text-encoder tuples should be `tuple()` and the
parent's length-equality validator should still pass — the porter
must zero out *all four* of `text_encoder_configs`,
`text_encoder_precisions`, `preprocess_text_funcs`,
`postprocess_text_funcs` together (we've seen this break before).
- [ ] **Component-bucket inheritance** (REVIEW item 24) — if the new
config goes under `vaes/`, it must subclass `VAEConfig` not
`EncoderConfig`. The bucket directory determines the base.
- [ ] **`__post_init__`** is used only to flip `load_encoder` /
`load_decoder` / etc. on the child configs, not to do any heavy
build.
#### Pipeline class review (`fastvideo/pipelines/basic/<family>/<family>_pipeline.py`)
- [ ] **`EntryClass` is a single class**, not a list.
- [ ] **`_required_config_modules`** lists exactly what the loader
reads from `model_index.json`. Missing keys cause silent loading
degradation.
- [ ] **`load_modules()`** does not have any `from diffusers import`
/ `from transformers import <ModelClass>` for production
components. If the porter explicitly opted into a temporary
diffusers shim, REJECT — the right move is to ship the native port
or hold the pipeline back (REVIEW item 30).
- [ ] **`torch.backends.*` flags** (TF32, cuDNN benchmark) — if set,
they're set **once in `load_modules`**, not per-call inside a stage
(REVIEW item 33). Mid-run flips invalidate the cuDNN algorithm
cache and amplify A2A SDE drift.
- [ ] **`create_pipeline_stages()`** uses standard stages from
`fastvideo.pipelines.stages` where possible, only subclassing when
the math diverges. New stage classes live in
`fastvideo/pipelines/basic/<family>/stages/` and are re-exported in
`fastvideo/pipelines/stages/__init__.py`.
- [ ] **One pipeline class for kwargs-driven variants** (REVIEW item
34) — T2A/A2A/inpaint that share weights/components shouldn't be
split into three classes. Triggers to split: separate
`_required_config_modules`, separate HF repo, divergent forward
signatures, separate `WorkloadType`. Reject the split unless one
of those applies.
#### Stages review (`fastvideo/pipelines/basic/<family>/stages/*.py`)
- [ ] **No `init_audio_strength = 0` / "0 or default" footguns** —
Python truthy-or fallbacks (`x or default`) silently swallow `0`,
`0.0`, `""`. Use `x if x is not None else default`.
- [ ] **Loud-fail on malformed kwarg combos** — e.g. inpaint without
mask should raise `ValueError`, not silently fall through to T2A.
- [ ] **Hot-path discipline** (REVIEW item 33): no per-step
`torch.zeros_like(...)` or `torch.randn_like(...)` inside the
sampler loop; pre-allocate buffers + reuse via `.normal_()` /
`.copy_()`.
- [ ] **No dead `batch.extra` writes** — every key written should be
read by a downstream stage.
- [ ] **Stage docstring** is one or two lines. No "Mirrors upstream X"
/ "Vendored from Y" provenance (REVIEW item 36); upstream
comparison belongs in the parity test, not the production docstring.
#### Presets + registry review
- [ ] **Sampling defaults** match the published model card example
block. Track the *model* defaults, not the upstream library's
*generic* defaults (this distinction shipped a 100% drift before
for Stable Audio).
- [ ] **`workload_type`** matches what the model actually does. If
the family is audio (T2A/A2A/AV) and `WorkloadType` doesn't yet
have those values (REVIEW item 28), accept `"t2v"` as a placeholder
but require a TODO comment + the porter to file a follow-up.
- [ ] **Registry detector** matches both the HF path and the
pipeline class name (`_class_name` from `model_index.json`).
- [ ] **`ALL_PRESETS`** export exists and is added to
`_register_presets()`'s group tuple in `fastvideo/registry.py`.
#### Tests review — **the most load-bearing review surface**
REVIEW item 16 is the single most expensive failure mode. A skipped
test reads as green in CI. Confirm explicitly:
- [ ] **Component parity tests exist** for every non-reused
component: DiT, VAE (if new), encoder (if new). Find them under
`tests/local_tests/<bucket>/test_<family>_*.py`.
- [ ] **Each component parity test produces a non-skip pass on the
reviewer's machine** if at all possible. If the porter says "I ran
it locally", ask for the diff numbers in the PR description.
- [ ] **Pipeline parity test exists** at
`tests/local_tests/pipelines/test_<family>_pipeline_parity.py` AND
has a non-skip pass — *not* a "skipped because the official clone
is missing" pass.
- [ ] **Smoke test exists** at `test_<family>_pipeline_smoke.py` — no
GPU, just import + registry + preset wiring. CI can run this even
if local-only parity tests skip.
- [ ] **Parity reference is the official upstream**, not diffusers
(REVIEW item 30). Diffusers parity is acceptable as a *secondary*
test only when (a) the published weights load through both and (b)
the official repo is also imported and compared.
- [ ] **Tolerances are scope-appropriate** (REVIEW item 21): single-
block + single-kernel = `atol=1e-4`; full-DiT cross-kernel = `0.1`
with a complementary `abs_mean drift < 5%` check; bare `assert_close`
alone with a loose tolerance is a smell.
- [ ] **Stub helpers for upstream private DSL deps** (REVIEW items
17, 18, 32) — if the upstream has `magi_compiler`-style imports
that aren't on PyPI, a small `tests/local_tests/helpers/<family>_upstream.py`
shim is fine. But if the deps it bypasses become real installs,
the shim must be deleted (item 32 — no zombie no-op shims).
- [ ] **GQA-aware kernel routing in stub paths** (REVIEW item 19) —
if the parity test routes upstream's flash_attn through SDPA, KV
heads must be `repeat_interleave`'d explicitly.
- [ ] **VAE normalization symmetry** (REVIEW item 20) — if upstream
bundles `z = z*std + mean` in `decode()` and FastVideo expects
pre-denormalized, the test must compensate explicitly.
#### Example script review (`examples/inference/basic/basic_<family>*.py`)
The bar here is high — the example is the user's entry point, not the
porter's debugging script.
- [ ] **User-story-shaped docstring** (REVIEW item 31) — at least one
`User story (<persona>):` block, then a "How it works" / "Picking
the dial" / "Tunable knobs" section. Not a code-narration docstring.
- [ ] **5-15 LOC of constants + one `generate_video()` call**
(REVIEW item 35). If the example carries a 25-line `_load_reference`
/ shape-norm / resample helper, that glue belongs *in the
pipeline*, not the example.
- [ ] **Pipeline accepts file paths** for any media-input kwarg
(`init_audio`, `inpaint_audio`, `image_path`). The example just
passes the path; pipeline does decode + resample internally.
- [ ] **No `torchaudio.load` on container formats** (mp4, m4a) —
routes through `torchcodec` → CUDA NVRTC. PyAV (already a
FastVideo dep, used by `_mux_audio`) handles all formats.
- [ ] **Tunable-knob defaults** match the model's published sweet
spot, not the upstream library's generic defaults.
#### Local repro doc review (optional but recommended)
If the family ships a `tests/local_tests/<family>.md`:
- [ ] **Setup section** covers HF gated access, optional inference
deps, upstream clone instructions, model cache pre-warm.
- [ ] **Per-test table** explains what each test compares against
with expected drift numbers.
- [ ] **Troubleshooting section** covers gated-skip behavior, batch-
vs-single-run flag interactions (cuDNN benchmark!), first-call
cache download blowup.
### 4. Cross-check the "first-class component" rule
Run `grep -rn "from diffusers import\|from transformers import"
fastvideo/pipelines/basic/<family>/ fastvideo/models/dits/<family>*
fastvideo/models/vaes/<family>* fastvideo/models/encoders/<family>*`.
The **only** acceptable hits are `from transformers import
<TokenizerFast>` (data-utility, no weights) and `from transformers
import T5EncoderModel` *only if* (a) it's loaded via HF and (b) the
T5 weights are absent from the model's checkpoint by design (e.g.
SA conditioner).
Anything else — `StableAudioDiTModel`, `AutoencoderKL`,
`UnetXxx`, `T2VPipeline` — is a REVIEW item 30 violation. Block
the PR with a pointer to that item.
### 5. Run the parity tests yourself if budget permits
The porter said it passes. Confirm:
```bash
# DiT/VAE/encoder component parity
pytest tests/local_tests/<bucket>/test_<family>_*.py -v -s
# Pipeline parity (the gate per add-model step 13(a))
pytest tests/local_tests/pipelines/test_<family>_pipeline_parity.py -v -s
# Smoke (no GPU)
pytest tests/local_tests/pipelines/test_<family>_pipeline_smoke.py -v
```
If any *parity* test SKIPs rather than PASSes on your machine and
you have the prerequisites set up, that means the porter never
actually verified parity (REVIEW item 16). Block.
### 6. Read REVIEW.md for any new failure modes the porter didn't address
The PR may also amend REVIEW.md with newly-discovered failure modes.
Skim those — they're typically the most accurate signal of what to
look for in *this specific* port. If the porter added a REVIEW item
and the linked code in the PR doesn't actually mitigate that item,
that's a contradiction — flag it.
### 7. Write the verdict
Structure your review comment as three buckets:
- **Block (must fix before merge)** — REVIEW item 30 violations,
skipped parity tests, missing required Files-table rows, dead
`batch.extra` writes that hide caller bugs, footguns like `0 or
default`.
- **Nit (would be better)** — naming convention drift, narrative
comments per REVIEW item 36, missing user-story docstrings,
hot-path allocations that don't change correctness.
- **Follow-up (not this PR)** — extracting test boilerplate to a
shared helper, future variant pipelines, performance benchmarks.
Always link each item back to the `add-model/REVIEW.md` item number
when applicable; that's the institutional memory and lets the porter
fix the issue with full context.
## Pitfall map (REVIEW.md items by where they show up in a diff)
Use this when you've spotted something off and want to find the
documented failure mode it corresponds to.
| If the diff has... | Suspect REVIEW item(s) |
|---|---|
| `from diffusers import <ModelClass>` in `fastvideo/...` | **30** (hard ban) |
| Raw `nn.Linear` in DiT projections | 10, 11 (only acceptable for MoE-packed weights with comment) |
| Custom RoPE that's *not* `_apply_rotary_emb` | Probably fine if the convention differs (e.g. halves-swap), but require a one-line WHY comment + a parity-test row |
| `parity_test.py` that calls `pytest.skip` unconditionally | **16** (silent no-op trap) — block |
| Component parity tests missing | **16** — block |
| New pipeline file with no `register_configs` call in registry.py | Files-table row 11 missing |
| `_required_config_modules` that doesn't match `model_index.json` keys | Silent loading degradation; block |
| `text_encoder_configs=tuple()` but other text-encoder tuples non-empty | Length-equality validator will fail at runtime |
| ArchConfig has `num_inference_steps` / `guidance_scale` / `flow_shift` | **15a** — pipeline-level fields leaking into ArchConfig; block |
| `<family>vae.py` for an arch-shared VAE (e.g. Oobleck) | **29** — should be `<arch>.py` (e.g. `oobleck.py`) |
| Pipeline imports `from transformers import T5EncoderModel` | OK only if the conditioner intentionally hides T5 from `named_parameters()` and the checkpoint omits T5 keys; document the exception |
| `torch.backends.*` set in a stage's `forward()` | **33** — should be one-shot in `load_modules` |
| `init_X = 0` silently treated as default via `or` | **33-adjacent** — Python truthy footgun; flag |
| Per-step `torch.zeros_like` / `randn_like` in sampler callback | **33** — pre-allocate; can also fix accuracy regressions |
| Magic `_DOWNSAMPLING_RATIO = 2048` in stage code | **33** — derive from the VAE config |
| `examples/.../basic_<family>*.py` >30 LOC | **35** — decode/resample/shape-norm belongs in the pipeline |
| Example uses `torchaudio.load(...)` on mp4 | **35** — torchcodec → NVRTC dep chain; use PyAV |
| Example docstring narrates code | **31** — needs `User story (<persona>):` block |
| Pipeline class docstring narrates upstream provenance | **36** — strip "Vendored from X" / "Mirrors upstream Y" |
| Comment says "previously this was X, we moved it because Y" | **36** — strip; goes in commit message, not code |
| Diff adds `tests/local_tests/helpers/<family>_upstream.py` that's a no-op | **32** — delete the no-op shim + its call sites |
| Diff adds `init_audio` / `inpaint_audio` / similar to `SamplingParam` | OK; verify they default to `None` and other families ignore |
| Multiple pipeline classes that share weights/components | **34** — should be one class with kwargs-driven modes |
| Config under `fastvideo/configs/models/encoders/` for a VAE | **24** — wrong bucket; should be under `vaes/` with `VAEConfig` base |
| New pipeline that depends on a not-yet-ported component (placeholder import) | **30** — hold the pipeline back until the component is ported |
## Outputs
A structured review comment on the PR with:
1. **Verdict** — `approve` / `request-changes` / `comment-only`.
2. **Block list** — REVIEW-item-linked issues that must be fixed.
3. **Nit list** — style / convention drift.
4. **Follow-up list** — items the porter should know about but
shouldn't fix in this PR.
5. **Parity numbers you confirmed** — diff_max / diff_mean / drift /
element-wise bound, per parity test you ran. Lets the next reviewer
skip re-running.
## Example review comment skeleton
```
## Review summary
**Verdict:** request-changes
I ran the smoke + DiT-component parity tests and read the full diff.
The native ports look right; two REVIEW-30 violations in the pipeline
file need resolving before merge, plus a few smaller nits.
### Block (must fix)
- `fastvideo/pipelines/basic/<family>/<family>_pipeline.py:NN` — `from
diffusers import <ModelClass>` violates REVIEW item 30 (hard ban).
Either ship the first-class port now or hold this pipeline back
until the component is ported.
- `tests/local_tests/pipelines/test_<family>_pipeline_parity.py` skips
on my machine even with HF token set (the upstream-clone path check
fails). REVIEW item 16: a parity test that always skips is worse
than no test. Update the path resolution or document it in
`tests/local_tests/<family>.md`.
### Nit
- `examples/inference/basic/basic_<family>.py` docstring is code-narration;
REVIEW item 31 wants a `User story (<persona>):` block.
- `<family>_pipeline.py:NN` has a per-call `torch.backends.cuda.matmul.allow_tf32 = False`;
REVIEW item 33 says move to one-shot in `load_modules`.
### Follow-up
- HF-token boilerplate is duplicated across 4 parity files. REVIEW
item 1-tier work, but doesn't need to land in this PR.
### Parity numbers I confirmed locally
| Test | diff_max | diff_mean | drift |
|---|---|---|---|
| DiT component parity | 0.0 | 0.0 | bit-identical |
| VAE decode parity | 0.0 | 0.0 | bit-identical |
| Pipeline parity (T2V, 25 steps) | 0.012 | 0.0009 | 0.31% |
```
## References
- `.agents/skills/add-model/SKILL.md` — the canonical procedure;
every "should be there" claim in this skill is a row in that
skill's Files-table.
- `.agents/skills/add-model/REVIEW.md` — 36 documented failure modes;
the Pitfall map above indexes them by where they show up in a diff.
- `.agents/skills/seed-ssim-references/SKILL.md` — for SSIM regression
add-on PRs (separate from this skill's scope).
@@ -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. |
+186 -2
View File
@@ -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
View File
@@ -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
+80
View File
@@ -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`,
});
}
+7 -4
View File
@@ -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"
+53 -12
View File
@@ -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 }}
+3 -3
View File
@@ -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" \
+1
View File
@@ -85,6 +85,7 @@ docs/distillation/examples/
dmd_t2v_output/
preprocess_output_text/
# Next.js / Node artifacts under ui/: see ui/.gitignore
.claude/
.codex/
+1
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
3.12
+2 -2
View File
@@ -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
+31 -6
View File
@@ -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 |
+5 -4
View File
@@ -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`
+10 -5
View File
@@ -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
+6 -5
View File
@@ -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
+177
View File
@@ -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`.
+24 -1
View File
@@ -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.
+1 -1
View File
@@ -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
+10 -15
View File
@@ -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
View File
@@ -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
```
+27 -19
View File
@@ -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
+129
View File
@@ -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)
+6
View File
@@ -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
+5
View File
@@ -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 -1
View File
@@ -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():
+1 -1
View File
@@ -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()
+109
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
def main():
@@ -5,7 +5,7 @@ import time
import gradio as gr
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
from copy import deepcopy
@@ -9,7 +9,7 @@ import tempfile
import gradio as gr
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
MODEL_PATH_MAPPING = {
@@ -185,7 +185,7 @@ class BaseModelDeployment:
def _initialize_generator(self, config: Dict[str, Any]) -> None:
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
print(f"Initializing model: {self.model_path}")
self.generator = VideoGenerator.from_pretrained(
@@ -1,5 +1,5 @@
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "./lora_out"
def main():
@@ -2,7 +2,7 @@
Inference using a LoRA checkpoint from FastVideo trainer.
"""
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "./lora_out"
def main():
+39 -1
View File
@@ -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)
+1 -1
View File
@@ -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 -1
View File
@@ -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__
+97
View File
@@ -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",
]
+623
View File
@@ -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",
]
+16
View File
@@ -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
+110
View File
@@ -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"]
+324
View File
@@ -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",
]
+261
View File
@@ -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",
]
+233
View File
@@ -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",
]
+101
View File
@@ -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:
+296
View File
@@ -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",
]
+738
View File
@@ -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
+188
View File
@@ -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",
]
+14
View File
@@ -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.
"""
+1 -1
View File
@@ -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)
+171
View File
@@ -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
+27 -2
View File
@@ -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, ))
-13
View File
@@ -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",
]
-18
View File
@@ -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
-23
View File
@@ -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
-21
View File
@@ -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
-55
View File
@@ -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
-25
View File
@@ -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 = ""
-20
View File
@@ -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
-69
View File
@@ -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
-25
View File
@@ -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
-154
View File
@@ -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
+1 -1
View File
@@ -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
+25 -69
View File
@@ -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",
]
+8 -6
View File
@@ -1,6 +1,5 @@
# 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
@@ -27,14 +26,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__":
+47 -69
View File
@@ -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)
+38 -2
View File
@@ -8,6 +8,8 @@ import uvicorn
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from fastvideo.api.presets import validate_preset_selection
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.state import (
DEFAULT_OUTPUT_DIR,
clear_state,
@@ -16,6 +18,7 @@ from fastvideo.entrypoints.openai.state import (
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.registry import get_preset_selection
logger = init_logger(__name__)
@@ -23,17 +26,40 @@ DEFAULT_HOST = "0.0.0.0"
DEFAULT_PORT = 8000
def _validate_default_request_against_preset(
default_request: GenerationRequest,
model_path: str,
) -> None:
"""Validate ``default_request.stage_overrides`` against the model's preset.
Called once at server startup from :func:`run_server`. The
``default_request`` is static server config, so validation results are
invariant across requests — there's no reason to re-run per request.
"""
if not default_request.stage_overrides:
return
preset_name, model_family = get_preset_selection(model_path)
if preset_name is None or model_family is None:
return
validate_preset_selection(
preset_name,
model_family,
stage_overrides=default_request.stage_overrides,
)
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
"""Load model on startup, clean up on shutdown"""
args: FastVideoArgs = app.state.fastvideo_args
output_dir: str = app.state.output_dir
default_request: GenerationRequest | None = getattr(app.state, "default_request", None)
logger.info("Loading model from %s ...", args.model_path)
generator = VideoGenerator.from_fastvideo_args(args)
logger.info("Model loaded successfully.")
set_state(generator, args, output_dir)
set_state(generator, args, output_dir, default_request=default_request)
yield # server is running
@@ -46,6 +72,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
def create_app(
fastvideo_args: FastVideoArgs,
output_dir: str = DEFAULT_OUTPUT_DIR,
default_request: GenerationRequest | None = None,
) -> FastAPI:
"""Build the FastAPI application with all routers mounted"""
@@ -56,6 +83,7 @@ def create_app(
)
app.state.fastvideo_args = fastvideo_args
app.state.output_dir = output_dir
app.state.default_request = default_request
app.add_middleware(
CORSMiddleware,
@@ -108,9 +136,17 @@ def run_server(
host: str = DEFAULT_HOST,
port: int = DEFAULT_PORT,
output_dir: str = DEFAULT_OUTPUT_DIR,
default_request: GenerationRequest | None = None,
):
"""Create the app and run it with uvicorn"""
app = create_app(fastvideo_args, output_dir=output_dir)
if default_request is not None:
_validate_default_request_against_preset(default_request, fastvideo_args.model_path)
app = create_app(
fastvideo_args,
output_dir=output_dir,
default_request=default_request,
)
logger.info("Starting FastVideo server on %s:%d", host, port)
logger.info("Model: %s", fastvideo_args.model_path)
+12 -2
View File
@@ -10,6 +10,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import FastVideoArgs
@@ -18,6 +19,7 @@ DEFAULT_OUTPUT_DIR = "outputs"
_generator: VideoGenerator | None = None
_fastvideo_args: FastVideoArgs | None = None
_output_dir: str = DEFAULT_OUTPUT_DIR
_default_request: GenerationRequest | None = None
def get_generator() -> VideoGenerator:
@@ -37,20 +39,28 @@ def get_output_dir() -> str:
return _output_dir
def get_default_request() -> GenerationRequest | None:
"""Return the ServeConfig.default_request set at startup, if any."""
return _default_request
def set_state(
generator: VideoGenerator,
fastvideo_args: FastVideoArgs,
output_dir: str,
default_request: GenerationRequest | None = None,
) -> None:
"""Set all server state at once (called from lifespan)."""
global _generator, _fastvideo_args, _output_dir
global _generator, _fastvideo_args, _output_dir, _default_request
_generator = generator
_fastvideo_args = fastvideo_args
_output_dir = output_dir
_default_request = default_request
def clear_state() -> None:
"""Clear server state on shutdown."""
global _generator, _fastvideo_args
global _generator, _fastvideo_args, _default_request
_generator = None
_fastvideo_args = None
_default_request = None
+53 -21
View File
@@ -19,7 +19,10 @@ from fastapi import (
)
from fastapi.responses import FileResponse
from fastvideo.api.compat import explicit_request_updates
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.state import (
get_default_request,
get_generator,
get_output_dir,
get_server_args,
@@ -42,49 +45,73 @@ logger = init_logger(__name__)
router = APIRouter(prefix="/v1/videos", tags=["videos"])
def _build_generation_kwargs(request_id: str, req: VideoGenerationsRequest) -> dict[str, Any]:
def _build_generation_kwargs(
request_id: str,
req: VideoGenerationsRequest,
default_request: GenerationRequest | None = None,
) -> dict[str, Any]:
"""Build a flat kwargs dict for ``generator.generate_video``.
Precedence (highest to lowest):
1. Request body — only fields the client explicitly sent
(``req.model_fields_set``, Pydantic v2).
2. ``default_request`` — only fields the operator explicitly set in
the serve YAML, projected via ``explicit_request_updates``. Schema
defaults on the dataclass are *not* treated as defaults here.
3. Hardcoded fallback (e.g. ``fps=24`` when neither side set it).
Why gate on ``model_fields_set`` / explicit paths? Both the request
Pydantic model and the ``GenerationRequest`` dataclass carry schema
defaults (e.g. ``seed=1024``, ``num_frames=125``). Without the gate
those would masquerade as intent and shadow the other side — the
gate preserves "operator pinned it" vs. "dataclass happened to have
that default."
"""
kwargs: dict[str, Any] = {}
if default_request is not None:
kwargs.update(explicit_request_updates(default_request))
body_set = req.model_fields_set
kwargs["prompt"] = req.prompt
# Resolution
if req.size:
if "size" in body_set and req.size:
w, h = parse_size(req.size)
if w is not None and h is not None:
kwargs["width"] = w
kwargs["height"] = h
# Frame count / duration
fps = req.fps if req.fps is not None else 24
kwargs["fps"] = fps
if "fps" in body_set and req.fps is not None:
kwargs["fps"] = req.fps
if req.num_frames is not None:
if "num_frames" in body_set and req.num_frames is not None:
kwargs["num_frames"] = req.num_frames
elif req.seconds is not None:
elif "seconds" in body_set and req.seconds is not None:
fps = kwargs.get("fps", 24)
kwargs["num_frames"] = fps * req.seconds
# Sampling parameters
if req.seed is not None:
if "seed" in body_set and req.seed is not None:
kwargs["seed"] = req.seed
if req.num_inference_steps is not None:
if ("num_inference_steps" in body_set and req.num_inference_steps is not None):
kwargs["num_inference_steps"] = req.num_inference_steps
if req.guidance_scale is not None:
if "guidance_scale" in body_set and req.guidance_scale is not None:
kwargs["guidance_scale"] = req.guidance_scale
if req.guidance_scale_2 is not None:
if "guidance_scale_2" in body_set and req.guidance_scale_2 is not None:
kwargs["guidance_scale_2"] = req.guidance_scale_2
if req.negative_prompt is not None:
if "negative_prompt" in body_set and req.negative_prompt is not None:
kwargs["negative_prompt"] = req.negative_prompt
if req.enable_teacache:
if "enable_teacache" in body_set and req.enable_teacache:
kwargs["enable_teacache"] = True
if req.true_cfg_scale is not None:
if "true_cfg_scale" in body_set and req.true_cfg_scale is not None:
kwargs["true_cfg_scale"] = req.true_cfg_scale
# Image-to-video input
if req.input_reference is not None:
if "input_reference" in body_set and req.input_reference is not None:
kwargs["image_path"] = req.input_reference
# Output path
output_dir = req.output_path or os.path.join(get_output_dir(), "videos")
kwargs.setdefault("fps", 24)
default_output_path = kwargs.pop("output_path", None)
body_output_dir = req.output_path if "output_path" in body_set else None
output_dir = body_output_dir or default_output_path or os.path.join(get_output_dir(), "videos")
os.makedirs(output_dir, exist_ok=True)
kwargs["output_path"] = os.path.join(output_dir, f"{request_id}.mp4")
kwargs["save_video"] = True
@@ -272,7 +299,12 @@ async def create_video(
logger.info("Video generation request %s: prompt=%s", request_id, req.prompt[:100])
gen_kwargs = _build_generation_kwargs(request_id, req)
# default_request was validated at server startup (run_server) and is
# read-only on the request hot path — _build_generation_kwargs and
# explicit_request_updates only read, so no per-request deepcopy needed.
default_request = get_default_request()
gen_kwargs = _build_generation_kwargs(request_id, req, default_request=default_request)
job = _make_video_job(request_id, req, gen_kwargs)
await VIDEO_STORE.upsert(request_id, job)
@@ -0,0 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.entrypoints.streaming.server import build_app, run_server
from fastvideo.entrypoints.streaming.session import (
Session,
SessionManager,
SessionState,
)
from fastvideo.entrypoints.streaming.session_store import (
BlobStore,
InMemoryBlobStore,
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
)
__all__ = [
"BlobStore",
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
"InMemoryBlobStore",
"InMemorySessionStore",
"Session",
"SessionManager",
"SessionState",
"SessionStore",
"build_app",
"run_server",
]
+252
View File
@@ -0,0 +1,252 @@
# SPDX-License-Identifier: Apache-2.0
"""JSON WebSocket protocol schemas for the streaming server.
Every control message shares the envelope ``{"type": <str>, ...}``.
Pydantic models live here so the server can parse / validate incoming
frames and emit well-typed outgoing frames without hand-rolled dicts.
The message catalogue matches the contract in
``docs/design/server_contracts/streaming.md``; additions must land in
both places in the same PR.
"""
from __future__ import annotations
from typing import Annotated, Any, Literal, Union
from pydantic import BaseModel, ConfigDict, Field
# ---------------------------------------------------------------------------
# Client → server
# ---------------------------------------------------------------------------
class SessionInitV2(BaseModel):
"""Opening frame the client sends after the WebSocket handshake."""
model_config = ConfigDict(extra="allow")
type: Literal["session_init_v2"]
client_id: str | None = None
preset: str | None = None
preset_label: str | None = None
curated_prompts: list[str] = Field(default_factory=list)
initial_image: dict[str, Any] | None = None
enhancement_enabled: bool = False
auto_extension_enabled: bool = False
loop_generation_enabled: bool = False
single_clip_mode: bool = False
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
continuation_state: dict[str, Any] | None = None
"""Optional ``{kind, payload}`` dict; hydrated into
:class:`fastvideo.api.ContinuationState` server-side."""
class SegmentPromptSource(BaseModel):
"""Request a new segment using a specific prompt."""
type: Literal["segment_prompt_source"]
prompt: str
negative_prompt: str | None = None
source: Literal["curated", "enhanced", "user", "auto_extension"] = "user"
seed: int | None = None
num_inference_steps: int | None = None
guidance_scale: float | None = None
class SeedPromptsUpdated(BaseModel):
type: Literal["seed_prompts_updated"]
seed_prompts: list[str] = Field(default_factory=list)
class EnhancementUpdated(BaseModel):
type: Literal["enhancement_updated"]
enabled: bool
class AutoExtensionUpdated(BaseModel):
type: Literal["auto_extension_updated"]
enabled: bool
class LoopGenerationUpdated(BaseModel):
type: Literal["loop_generation_updated"]
enabled: bool
class GenerationPausedUpdated(BaseModel):
type: Literal["generation_paused_updated"]
paused: bool
class SnapshotState(BaseModel):
"""Request the current ``ContinuationState`` for export."""
type: Literal["snapshot_state"]
ClientMessage = Annotated[
Union[ # noqa: UP007 - Annotated requires Union for discriminator
SessionInitV2,
SegmentPromptSource,
SeedPromptsUpdated,
EnhancementUpdated,
AutoExtensionUpdated,
LoopGenerationUpdated,
GenerationPausedUpdated,
SnapshotState,
],
Field(discriminator="type"),
]
# ---------------------------------------------------------------------------
# Server → client
# ---------------------------------------------------------------------------
class QueueStatus(BaseModel):
type: Literal["queue_status"] = "queue_status"
position: int
queue_depth: int
class GpuAssigned(BaseModel):
type: Literal["gpu_assigned"] = "gpu_assigned"
gpu_id: int
session_timeout: int
class Ltx2StreamStart(BaseModel):
type: Literal["ltx2_stream_start"] = "ltx2_stream_start"
preset: str | None = None
width: int
height: int
fps: int
num_frames: int
class Ltx2SegmentStart(BaseModel):
type: Literal["ltx2_segment_start"] = "ltx2_segment_start"
segment_idx: int
prompt: str
total_steps: int
class StepComplete(BaseModel):
type: Literal["step_complete"] = "step_complete"
segment_idx: int
step: int
total_steps: int
stage: str = "denoise"
class MediaInit(BaseModel):
"""Descriptor for the fMP4 initialization segment that follows."""
type: Literal["media_init"] = "media_init"
segment_idx: int
mime: str = "video/mp4; codecs=\"avc1.64001f, mp4a.40.2\""
stream_id: str
mode: Literal["av_fmp4"] = "av_fmp4"
class MediaSegmentComplete(BaseModel):
type: Literal["media_segment_complete"] = "media_segment_complete"
segment_idx: int
stream_id: str
chunks: int
duration_ms: float | None = None
pts_base_ms: float | None = None
class Ltx2SegmentComplete(BaseModel):
type: Literal["ltx2_segment_complete"] = "ltx2_segment_complete"
segment_idx: int
generation_time_ms: float
e2e_latency_ms: float | None = None
class Ltx2StreamComplete(BaseModel):
type: Literal["ltx2_stream_complete"] = "ltx2_stream_complete"
reason: Literal["segment_cap", "stop_requested", "error"] = "stop_requested"
class SessionTimeout(BaseModel):
type: Literal["session_timeout"] = "session_timeout"
timeout_seconds: int
class ContinuationStateSnapshot(BaseModel):
type: Literal["continuation_state_snapshot"] = "continuation_state_snapshot"
state: dict[str, Any]
"""``{kind, payload}`` dict matching
:class:`fastvideo.api.ContinuationState`."""
class ErrorMessage(BaseModel):
type: Literal["error"] = "error"
code: Literal[
"session_rejected",
"invalid_message",
"preset_mismatch",
"gpu_unavailable",
"worker_failed",
"upstream_timeout",
"internal_error",
] = "internal_error"
message: str
retryable: bool = False
ServerMessage = Union[ # noqa: UP007 - pydantic Union handling
QueueStatus,
GpuAssigned,
Ltx2StreamStart,
Ltx2SegmentStart,
StepComplete,
MediaInit,
MediaSegmentComplete,
Ltx2SegmentComplete,
Ltx2StreamComplete,
SessionTimeout,
ContinuationStateSnapshot,
ErrorMessage,
]
def parse_client_message(raw: dict[str, Any]) -> ClientMessage:
"""Parse an incoming WebSocket dict into a typed client message.
Unknown ``type`` values raise :class:`pydantic.ValidationError`; the
server handler turns that into an ``error`` frame with
``code="invalid_message"``.
"""
from pydantic import TypeAdapter
return TypeAdapter(ClientMessage).validate_python(raw)
__all__ = [
"AutoExtensionUpdated",
"ClientMessage",
"ContinuationStateSnapshot",
"EnhancementUpdated",
"ErrorMessage",
"GenerationPausedUpdated",
"GpuAssigned",
"Ltx2SegmentComplete",
"Ltx2SegmentStart",
"Ltx2StreamComplete",
"Ltx2StreamStart",
"LoopGenerationUpdated",
"MediaInit",
"MediaSegmentComplete",
"QueueStatus",
"SeedPromptsUpdated",
"SegmentPromptSource",
"ServerMessage",
"SessionInitV2",
"SessionTimeout",
"SnapshotState",
"StepComplete",
"parse_client_message",
]
+531
View File
@@ -0,0 +1,531 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-generator FastAPI + WebSocket streaming server."""
from __future__ import annotations
import asyncio
import contextlib
import os
import time
from dataclasses import dataclass
from typing import Any, Protocol
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
InputConfig,
OutputConfig,
SamplingConfig,
ServeConfig,
)
from fastvideo.entrypoints.streaming.protocol import (
AutoExtensionUpdated,
ContinuationStateSnapshot,
EnhancementUpdated,
ErrorMessage,
GenerationPausedUpdated,
GpuAssigned,
LoopGenerationUpdated,
Ltx2SegmentComplete,
Ltx2SegmentStart,
Ltx2StreamComplete,
Ltx2StreamStart,
MediaInit,
MediaSegmentComplete,
QueueStatus,
SeedPromptsUpdated,
SegmentPromptSource,
SessionInitV2,
SnapshotState,
StepComplete,
parse_client_message,
)
from fastvideo.entrypoints.streaming.session import (
InvalidSessionTransition,
Session,
SessionManager,
SessionRejected,
SessionState,
)
from fastvideo.entrypoints.streaming.session_init_image import (
persist_session_init_image, )
from fastvideo.entrypoints.streaming.session_store import (
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# RFC 6455 WebSocket close codes used by the server.
_WS_CLOSE_UNSUPPORTED_DATA = 1003
_WS_CLOSE_TRY_AGAIN_LATER = 1013
class _GeneratorProto(Protocol):
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
def generate(self, request: GenerationRequest) -> Any:
...
@dataclass
class ServerState:
serve_config: ServeConfig
generator: _GeneratorProto
sessions: SessionManager
session_store: SessionStore
def build_app(
serve_config: ServeConfig,
generator: _GeneratorProto,
*,
session_store: SessionStore | None = None,
) -> FastAPI:
"""Build the FastAPI app used by :func:`run_server`.
Exposed so tests can drive the WebSocket endpoint in-process via
``starlette.testclient.TestClient(app).websocket_connect(...)``.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
"server; got None. Add a `streaming:` block to your serve config.")
sessions = SessionManager(
segment_cap=serve_config.streaming.generation_segment_cap,
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
)
state = ServerState(
serve_config=serve_config,
generator=generator,
sessions=sessions,
session_store=session_store or InMemorySessionStore(),
)
app = FastAPI(title="FastVideo Streaming")
@app.get("/health")
async def _health() -> JSONResponse:
return JSONResponse({
"status": "ok",
"sessions": len(state.sessions),
"stream_mode": state.serve_config.streaming.stream_mode,
})
@app.websocket("/v1/stream")
async def _stream(websocket: WebSocket) -> None:
await websocket.accept()
try:
session = state.sessions.create()
except SessionRejected as exc:
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
return
try:
await _handle_session(websocket, session, state)
except WebSocketDisconnect:
logger.info("session %s: client disconnected", session.id[:8])
except Exception: # pragma: no cover - defensive catch-all
logger.exception("session %s: unhandled error", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
finally:
_cleanup_session(session, state)
app.state.server_state = state
return app
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
"""Launch the streaming server.
Boots a :class:`fastvideo.VideoGenerator` from
``serve_config.generator`` unless ``generator`` is provided, then
serves ``build_app(...)`` via uvicorn.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
"got None. Add a `streaming:` block to your serve config.")
import uvicorn
if generator is None:
from fastvideo import VideoGenerator # lazy to avoid boot cost
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
app = build_app(serve_config, generator)
uvicorn.run(
app,
host=serve_config.server.host,
port=serve_config.server.port,
)
async def _handle_session(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
init = await _read_init_message(websocket, session, state)
if init is None:
return
await _apply_session_init(session, init, state)
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
session.transition(SessionState.GPU_BINDING)
await _send_json(websocket, GpuAssigned(
gpu_id=0,
session_timeout=state.sessions.session_timeout_seconds,
))
session.transition(SessionState.ACTIVE)
await _send_json(websocket, _build_stream_start(session, state))
try:
await _run_segment_loop(websocket, session, state)
finally:
with contextlib.suppress(RuntimeError):
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
async def _read_init_message(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> SessionInitV2 | None:
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: init timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return None
except WebSocketDisconnect:
return None
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
return None
if not isinstance(parsed, SessionInitV2):
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
return None
return parsed
async def _reject_init(
websocket: WebSocket,
session: Session,
message: str,
close_reason: str,
) -> None:
await _send_error(websocket, "invalid_message", message, retryable=False)
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.REJECTED)
async def _apply_session_init(
session: Session,
init: SessionInitV2,
state: ServerState,
) -> None:
session.client_id = init.client_id
session.preset = init.preset
session.preset_label = init.preset_label
session.curated_prompts = list(init.curated_prompts)
session.enhancement_enabled = init.enhancement_enabled
session.auto_extension_enabled = init.auto_extension_enabled
session.loop_generation_enabled = init.loop_generation_enabled
session.single_clip_mode = init.single_clip_mode
session.stream_mode = init.stream_mode
if init.initial_image is not None:
# Decode + disk write off the event loop; payload is up to 32 MiB.
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
if image is not None:
session.metadata["session_init_image"] = image.path
if init.continuation_state is not None:
session.continuation_state = _coerce_state(init.continuation_state)
if session.continuation_state is not None:
state.session_store.store(session.id, session.continuation_state)
async def _run_segment_loop(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
cap = state.sessions.segment_cap
while True:
if session.segment_cap_reached(cap):
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
return
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: idle timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return
except WebSocketDisconnect:
return
session.touch()
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
continue
if isinstance(parsed, SnapshotState):
snap = state.session_store.snapshot(session.id)
if snap is None:
await _send_error(websocket,
"internal_error",
"no continuation state available for session",
retryable=False)
continue
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
continue
if isinstance(parsed, SegmentPromptSource):
await _run_segment(websocket, session, state, parsed)
continue
# Silently ignore unknown-but-valid types (additive-evolution
# rule in streaming.md).
_apply_toggle(session, parsed)
async def _run_segment(
websocket: WebSocket,
session: Session,
state: ServerState,
message: SegmentPromptSource,
) -> None:
request = _build_generation_request(session, message, state)
segment_idx = session.segment_idx
await _send_json(
websocket,
Ltx2SegmentStart(
segment_idx=segment_idx,
prompt=message.prompt,
total_steps=request.sampling.num_inference_steps,
))
start = time.perf_counter()
loop = asyncio.get_running_loop()
# TODO: executor-wrapped generate() cannot be cancelled, so a
# client disconnect mid-segment leaves the GPU work running to
# completion. Real cancellation needs the generate_async API.
try:
result = await loop.run_in_executor(None, state.generator.generate, request)
except Exception as exc:
logger.exception("session %s: generator failed", session.id[:8])
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
elapsed_ms = (time.perf_counter() - start) * 1000.0
frames = _extract_frames(result)
if not frames:
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
# Synchronous generator call has no per-step hook; emit one
# terminal StepComplete so observability wiring still sees the
# segment finish.
total = request.sampling.num_inference_steps
await _send_json(websocket, StepComplete(
segment_idx=segment_idx,
step=total,
total_steps=total,
stage="denoise",
))
encoder = FragmentedMP4Encoder(
width=request.sampling.width,
height=request.sampling.height,
fps=request.sampling.fps,
segment_idx=segment_idx,
)
chunks_relayed = 0
async with encoder:
init_sent = False
async for chunk in encoder.encode(frames):
if chunk.kind == "init":
await _send_json(websocket, MediaInit(
segment_idx=segment_idx,
stream_id=chunk.stream_id,
))
init_sent = True
await websocket.send_bytes(chunk.data)
if init_sent and chunk.kind == "media":
chunks_relayed += 1
await _send_json(
websocket,
MediaSegmentComplete(
segment_idx=segment_idx,
stream_id=encoder.stream_id,
chunks=chunks_relayed,
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
))
new_state = _extract_state(result)
if new_state is not None:
session.continuation_state = new_state
state.session_store.store(session.id, new_state)
session.segment_idx += 1
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ACTIVE)
await _send_json(
websocket,
Ltx2SegmentComplete(
segment_idx=segment_idx,
generation_time_ms=elapsed_ms,
e2e_latency_ms=elapsed_ms,
))
def _build_stream_start(
session: Session,
state: ServerState,
) -> Ltx2StreamStart:
default = state.serve_config.default_request
return Ltx2StreamStart(
preset=session.preset,
width=default.sampling.width,
height=default.sampling.height,
fps=default.sampling.fps,
num_frames=default.sampling.num_frames,
)
def _build_generation_request(
session: Session,
message: SegmentPromptSource,
state: ServerState,
) -> GenerationRequest:
# Start from the operator-pinned default_request to pick up the
# preset-selected sampling knobs; override with per-message values.
base = state.serve_config.default_request
sampling_kwargs: dict[str, Any] = {
"num_videos_per_prompt":
base.sampling.num_videos_per_prompt,
"seed":
message.seed if message.seed is not None else base.sampling.seed,
"num_frames":
base.sampling.num_frames,
"height":
base.sampling.height,
"width":
base.sampling.width,
"fps":
base.sampling.fps,
"num_inference_steps":
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
"guidance_scale":
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
}
request = GenerationRequest(
prompt=message.prompt,
negative_prompt=message.negative_prompt or base.negative_prompt,
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
sampling=SamplingConfig(**sampling_kwargs),
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
state=session.continuation_state,
)
return request
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
kind = raw.get("kind")
payload = raw.get("payload")
if not isinstance(kind, str) or not isinstance(payload, dict):
return None
return ContinuationState(kind=kind, payload=payload)
def _apply_toggle(session: Session, message: Any) -> None:
if isinstance(message, EnhancementUpdated):
session.enhancement_enabled = message.enabled
elif isinstance(message, AutoExtensionUpdated):
session.auto_extension_enabled = message.enabled
elif isinstance(message, LoopGenerationUpdated):
session.loop_generation_enabled = message.enabled
elif isinstance(message, GenerationPausedUpdated):
session.generation_paused = message.paused
elif isinstance(message, SeedPromptsUpdated):
session.curated_prompts = list(message.seed_prompts)
def _extract_frames(result: Any) -> list:
if hasattr(result, "frames"):
return list(result.frames or [])
if isinstance(result, dict):
return list(result.get("frames") or [])
return []
def _extract_state(result: Any) -> ContinuationState | None:
state = getattr(result, "state", None)
if state is None and isinstance(result, dict):
state = result.get("state")
if isinstance(state, ContinuationState):
return state
if isinstance(state, dict):
return _coerce_state(state)
return None
async def _send_json(websocket: WebSocket, message: Any) -> None:
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
await websocket.send_json(payload)
async def _send_error(
websocket: WebSocket,
code: str,
message: str,
*,
retryable: bool,
) -> None:
await _send_json(
websocket,
ErrorMessage(code=code, message=message, retryable=retryable),
)
def _cleanup_session(session: Session, state: ServerState) -> None:
state.sessions.close(session.id)
state.session_store.drop(session.id)
init_image_path = session.metadata.get("session_init_image")
if isinstance(init_image_path, str):
with contextlib.suppress(FileNotFoundError):
os.unlink(init_image_path)
__all__ = [
"ServerState",
"build_app",
"run_server",
]

Some files were not shown because too many files have changed in this diff Show More