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
69 changed files with 9436 additions and 274 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. |
@@ -29,7 +29,7 @@ surfaces:
vae_cpu_offload: generator.engine.offload.vae
pin_cpu_memory: generator.engine.offload.pin_cpu_memory
enable_torch_compile: generator.engine.compile.enabled
torch_compile_kwargs: generator.engine.compile.kwargs
torch_compile_kwargs: generator.engine.compile.backend,fullgraph,mode,dynamic,extras
disable_autocast: generator.engine.disable_autocast
enable_stage_verification: generator.engine.enable_stage_verification
prompt_txt: request.inputs.prompt_path
@@ -40,8 +40,8 @@ surfaces:
init_weights_from_safetensors_2: generator.pipeline.components.transformer_2_weights
override_pipeline_cls_name: generator.pipeline.components.override_pipeline_cls_name
boundary_ratio: request.sampling.boundary_ratio
ltx2_vae_tiling: generator.pipeline.vae_tiling
preset_owned:
ltx2_vae_tiling: generator.pipeline.preset_overrides.ltx2.vae_tiling
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
@@ -354,6 +354,8 @@ surfaces:
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
+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`.
+22
View File
@@ -86,3 +86,25 @@ sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
- Learning rate: 2e-5
- Training steps: 3000 (~12 hours)
- HSDP shard dim: 1
## 🧭 Note on `real_score_guidance_scale`
The teacher CFG used inside the DMD loss follows the DMD2 reference
implementation and uses the parameterization
```
x = x_cond + w * (x_cond - x_uncond)
```
rather than the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)`. The
two are mathematically equivalent up to a constant offset:
| `real_score_guidance_scale` (`w`) | Equivalent standard CFG (`w + 1`) | Output |
|-----------------------------------|-----------------------------------|-----------------------|
| `-1` | `0` | unconditional |
| `0` | `1` | conditional |
| `3.5` (default) | `4.5` | strong guidance |
So `real_score_guidance_scale` should be read as the **extra** guidance
strength added on top of the conditional prediction. When porting values
from a paper that uses the Ho & Salimans form, subtract 1.
+10 -1
View File
@@ -61,7 +61,16 @@ has_cmake_arg() {
}
detect_with_torch() {
uv run --active --no-project python -c "import torch
# Prefer the active venv's python directly over `uv run --active --no-project`,
# which on some uv versions provisions its own interpreter and misses packages
# installed into VIRTUAL_ENV.
local py
if [[ -n "${VIRTUAL_ENV:-}" && -x "${VIRTUAL_ENV}/bin/python" ]]; then
py="${VIRTUAL_ENV}/bin/python"
else
py="$(command -v python3 || command -v python)"
fi
"${py}" -c "import torch
if not torch.cuda.is_available():
raise RuntimeError('torch.cuda.is_available() is false')
mj, mn = torch.cuda.get_device_capability(0)
@@ -5,6 +5,11 @@ from fastvideo_kernel.ops import (
video_sparse_attn,
)
from fastvideo_kernel.block_sparse_attn import (
block_sparse_attn,
block_sparse_attn_from_indices,
)
from fastvideo_kernel.vmoba import (
moba_attn_varlen,
process_moba_input,
@@ -22,6 +27,8 @@ from fastvideo_kernel.turbodiffusion_ops import (
__all__ = [
"sliding_tile_attention",
"video_sparse_attn",
"block_sparse_attn",
"block_sparse_attn_from_indices",
"moba_attn_varlen",
"process_moba_input",
"process_moba_output",
@@ -1,3 +1,5 @@
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
from __future__ import annotations
import os
@@ -6,6 +8,11 @@ from typing import Tuple
import torch
# ---------------------------------------------------------------------------
# Backend selection helpers
# ---------------------------------------------------------------------------
def _get_sm90_ops():
try:
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
@@ -25,38 +32,66 @@ def _is_sm90() -> bool:
def _force_triton() -> bool:
# Force Triton even on SM90 and even if the compiled extension is available.
# Useful for CI / debugging / parity testing.
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Preferred map->index conversion used by the wrapper.
# ---------------------------------------------------------------------------
# Index helpers
# ---------------------------------------------------------------------------
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
"""
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Compact a bool block_map to (q2k_idx, q2k_num). Legacy path only."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
if block_map.dim() != 4:
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
raise ValueError(
f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), "
f"got shape={tuple(block_map.shape)}"
)
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
if not block_map.is_cuda:
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
raise RuntimeError(
"block_map must be a CUDA tensor (Triton map_to_index required)."
)
try:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index
except Exception as e: # pragma: no cover - environment issue
raise ImportError(
"Triton map_to_index is required but not available. "
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
"Ensure Triton is installed and "
"fastvideo_kernel.triton_kernels.index is importable."
) from e
return triton_map_to_index(block_map)
def _invert_indices_for_backward(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
from fastvideo_kernel.triton_kernels.index import invert_indices
return invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
"""Return `t` as a contiguous int32 tensor, raising a clear error on CPU input."""
if not t.is_cuda:
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
if t.dtype != torch.int32:
t = t.to(torch.int32)
if not t.is_contiguous():
t = t.contiguous()
return t
# ---------------------------------------------------------------------------
# Triton backend custom ops (index-native)
# ---------------------------------------------------------------------------
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_triton",
mutates_args=(),
@@ -66,34 +101,40 @@ def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
triton_block_sparse_attn_forward,
)
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
o, M = triton_block_sparse_attn_forward(
q.contiguous(),
k.contiguous(),
v.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
M = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
device=q.device,
dtype=torch.float32,
)
return o, M
@@ -109,20 +150,32 @@ def block_sparse_attn_backward_triton(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
triton_block_sparse_attn_backward,
)
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
# o/M are kernel outputs so are already contiguous.
dq, dk, dv = triton_block_sparse_attn_backward(
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
grad_output.contiguous(),
q.contiguous(),
k.contiguous(),
v.contiguous(),
o,
M,
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
variable_block_sizes,
)
return dq, dk, dv
@@ -135,7 +188,8 @@ def _block_sparse_attn_backward_triton_fake(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q)
@@ -144,19 +198,28 @@ def _block_sparse_attn_backward_triton_fake(
return dq, dk, dv
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def _setup_context_triton(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
block_sparse_attn_triton.register_autograd(
_backward_triton, setup_context=_setup_context_triton
)
# ---------------------------------------------------------------------------
# SM90 backend custom ops (index-native)
# ---------------------------------------------------------------------------
@torch.library.custom_op(
@@ -168,21 +231,21 @@ def block_sparse_attn_sm90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
block_sparse_fwd, _ = _get_sm90_ops()
if block_sparse_fwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o_padded, lse_padded
@@ -192,11 +255,16 @@ def _block_sparse_attn_sm90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q_padded)
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
lse = torch.empty(
(q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1),
device=q_padded.device,
dtype=torch.float32,
)
return o, lse
@@ -212,30 +280,34 @@ def block_sparse_attn_backward_sm90(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
_, block_sparse_bwd = _get_sm90_ops()
if block_sparse_bwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
dq, dk, dv = block_sparse_bwd(
q_padded,
k_padded,
v_padded,
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
o_padded,
lse_padded,
grad_output_padded,
grad_output_padded.contiguous(),
k2q_idx,
k2q_num,
variable_block_sizes.int(),
variable_block_sizes,
)
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
# C++ kernel returns fp32 grads; cast back to the input dtype.
out_dtype = grad_output_padded.dtype
return dq.to(out_dtype), dk.to(out_dtype), dv.to(out_dtype)
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
@@ -246,7 +318,8 @@ def _block_sparse_attn_backward_sm90_fake(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q_padded)
@@ -255,21 +328,57 @@ def _block_sparse_attn_backward_sm90_fake(
return dq, dk, dv
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
)
return dq, dk, dv, None, None
def _setup_context_sm90(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, lse = output
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
block_sparse_attn_sm90.register_autograd(
_backward_sm90, setup_context=_setup_context_sm90
)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def block_sparse_attn_from_indices(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Block-sparse attention with autograd, taking compact per-row KV indices."""
# Normalize index tensors once at the public boundary so the custom ops
# and their fakes can assume int32/contiguous. No-op on well-formed input.
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
variable_block_sizes = _as_int32_contig(variable_block_sizes, "variable_block_sizes")
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
use_sm90 = (
(not _force_triton())
and _is_sm90()
and block_sparse_fwd is not None
and block_sparse_bwd is not None
)
if use_sm90:
return block_sparse_attn_sm90(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
def block_sparse_attn(
@@ -279,16 +388,8 @@ def block_sparse_attn(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Unified block-sparse attention op with autograd support.
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
"""
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
"""Bool-mask compat wrapper; prefer block_sparse_attn_from_indices."""
q2k_idx, q2k_num = _map_to_index(block_map)
return block_sparse_attn_from_indices(
q, k, v, q2k_idx, q2k_num, variable_block_sizes
)
@@ -1,6 +1,6 @@
import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
# Try to load the C++ extension
@@ -125,13 +125,18 @@ def video_sparse_attn(
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, q_seq_len, dim)
# Sparse branch
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
topk_idx = torch.topk(scores, topk, dim=-1).indices
mask = torch.zeros_like(scores,
dtype=torch.bool).scatter_(-1, topk_idx, True)
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
q2k_idx = topk_idx.to(torch.int32).contiguous()
q2k_num = torch.full(
(batch, heads, q_num_blocks),
topk,
dtype=torch.int32,
device=q.device,
)
out_s = block_sparse_attn_from_indices(
q, k, v, q2k_idx, q2k_num, variable_block_sizes
)[0]
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
@@ -1,9 +1,10 @@
## pytorch sdpa version of block sparse ##
from typing import Tuple
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
@@ -153,3 +154,114 @@ def map_to_index(block_map: torch.Tensor):
)
return index, index_num
@triton.jit
def _invert_indices_kernel(
q2k_idx_ptr,
q2k_num_ptr,
k2q_idx_ptr,
k2q_num_ptr,
q2k_idx_b, q2k_idx_h, q2k_idx_q, q2k_idx_k,
q2k_num_b, q2k_num_h, q2k_num_q,
k2q_idx_b, k2q_idx_h, k2q_idx_k, k2q_idx_q,
k2q_num_b, k2q_num_h, k2q_num_k,
MAX_KV_PER_Q: tl.constexpr,
):
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
pid_q = tl.program_id(2)
n = tl.load(
q2k_num_ptr
+ pid_b * q2k_num_b
+ pid_h * q2k_num_h
+ pid_q * q2k_num_q
)
q2k_row = (
q2k_idx_ptr
+ pid_b * q2k_idx_b
+ pid_h * q2k_idx_h
+ pid_q * q2k_idx_q
)
for i in tl.range(0, MAX_KV_PER_Q):
if i < n:
kv = tl.load(q2k_row + i * q2k_idx_k)
count_ptr = (
k2q_num_ptr
+ pid_b * k2q_num_b
+ pid_h * k2q_num_h
+ kv * k2q_num_k
)
pos = tl.atomic_add(count_ptr, 1)
tl.store(
k2q_idx_ptr
+ pid_b * k2q_idx_b
+ pid_h * k2q_idx_h
+ kv * k2q_idx_k
+ pos * k2q_idx_q,
pid_q,
)
def invert_indices(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
if q2k_idx.dim() != 4:
raise ValueError(
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
)
if q2k_num.dim() != 3:
raise ValueError(
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
)
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
raise RuntimeError("invert_indices requires CUDA tensors.")
B, H, Nq, Mk = q2k_idx.shape
if q2k_num.shape != (B, H, Nq):
raise ValueError(
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
f"[B, H, Nq] = {(B, H, Nq)}"
)
q2k_idx = q2k_idx.contiguous()
q2k_num = q2k_num.contiguous()
if q2k_idx.dtype != torch.int32:
q2k_idx = q2k_idx.to(torch.int32)
if q2k_num.dtype != torch.int32:
q2k_num = q2k_num.to(torch.int32)
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
# `Nq` is a tight upper bound on the compacted K->Q slots.
k2q_idx = torch.empty(
(B, H, num_kv_blocks, Nq),
dtype=torch.int32,
device=q2k_idx.device,
)
k2q_num = torch.zeros(
(B, H, num_kv_blocks),
dtype=torch.int32,
device=q2k_idx.device,
)
grid = (B, H, Nq)
_invert_indices_kernel[grid](
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
q2k_idx.stride(0), q2k_idx.stride(1), q2k_idx.stride(2), q2k_idx.stride(3),
q2k_num.stride(0), q2k_num.stride(1), q2k_num.stride(2),
k2q_idx.stride(0), k2q_idx.stride(1), k2q_idx.stride(2), k2q_idx.stride(3),
k2q_num.stride(0), k2q_num.stride(1), k2q_num.stride(2),
MAX_KV_PER_Q=Mk,
)
return k2q_idx, k2q_num
+118 -9
View File
@@ -16,6 +16,8 @@ from fastvideo.api.request_metadata import (
reset_tracking_roots,
)
from fastvideo.api.schema import (
CompileConfig,
ContinuationState,
GenerationRequest,
GeneratorConfig,
InputConfig,
@@ -25,6 +27,10 @@ from fastvideo.api.schema import (
)
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)}
@@ -38,6 +44,10 @@ _LEGACY_REQUEST_ALIASES = {
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
"embedded_cfg_scale",
})
# torch.compile kwargs that map to first-class CompileConfig fields.
_COMPILE_TYPED_KEYS = ("backend", "fullgraph", "mode", "dynamic")
# LTX-2 refine flat kwargs (init + per-request) known to FastVideoArgs.
_LTX2_REFINE_FLAT_KEYS = (refine_preset_override_fields() | refine_stage_override_fields())
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
@@ -80,6 +90,8 @@ def legacy_from_pretrained_to_config(
components: dict[str, Any] = {}
quantization: dict[str, Any] = {}
experimental: dict[str, Any] = {}
preset_overrides: dict[str, Any] = {}
preset_refine: dict[str, Any] = {}
for key, value in kwargs.items():
if key == "revision":
@@ -106,8 +118,33 @@ def legacy_from_pretrained_to_config(
offload["pin_cpu_memory"] = value
elif key == "enable_torch_compile":
compile_config["enabled"] = value
elif key == "enable_torch_compile_text_encoder":
compile_config["text_encoder_enabled"] = value
elif key == "torch_compile_kwargs":
compile_config["kwargs"] = deepcopy(value)
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
for first_class in _COMPILE_TYPED_KEYS:
if first_class in remaining:
compile_config[first_class] = remaining.pop(first_class)
if remaining:
compile_config["extras"] = remaining
elif key == "ltx2_vae_tiling":
pipeline["vae_tiling"] = value
elif key == "config_model_path":
components["config_root"] = value
elif key == "ltx2_refine_enabled":
preset_refine["enabled"] = value
elif key == "ltx2_refine_upsampler_path":
# Empty string means "no upsampler"; keep typed None.
components["upsampler_weights"] = value or None
elif key == "ltx2_refine_lora_path":
# Empty string means "no refine LoRA"; keep typed None.
components["lora_path"] = value or None
elif key == "ltx2_refine_add_noise":
preset_refine["add_noise"] = value
elif key == "ltx2_refine_num_inference_steps":
preset_refine["num_inference_steps"] = value
elif key == "ltx2_refine_guidance_scale":
preset_refine["guidance_scale"] = value
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
engine[key] = value
elif key == "override_text_encoder_quant":
@@ -147,6 +184,10 @@ def legacy_from_pretrained_to_config(
if components:
pipeline["components"] = components
if preset_refine:
preset_overrides["refine"] = preset_refine
if preset_overrides:
pipeline["preset_overrides"] = preset_overrides
if experimental:
pipeline["experimental"] = experimental
if pipeline:
@@ -162,12 +203,8 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
unsupported.append("pipeline.preset")
if normalized.pipeline.preset_version is not None:
unsupported.append("pipeline.preset_version")
if normalized.pipeline.components.config_root is not None:
unsupported.append("pipeline.components.config_root")
if normalized.pipeline.components.vae_weights is not None:
unsupported.append("pipeline.components.vae_weights")
if normalized.pipeline.components.upsampler_weights is not None:
unsupported.append("pipeline.components.upsampler_weights")
if unsupported:
joined = ", ".join(unsupported)
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
@@ -191,13 +228,21 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
"vae_cpu_offload": engine.offload.vae,
"pin_cpu_memory": engine.offload.pin_cpu_memory,
"enable_torch_compile": engine.compile.enabled,
"torch_compile_kwargs": deepcopy(engine.compile.kwargs),
"torch_compile_kwargs": _compile_config_to_torch_kwargs(engine.compile),
"enable_stage_verification": engine.enable_stage_verification,
"use_fsdp_inference": engine.use_fsdp_inference,
"disable_autocast": engine.disable_autocast,
}
if normalized.pipeline.workload_type is not None:
kwargs["workload_type"] = normalized.pipeline.workload_type
if normalized.pipeline.vae_tiling is not None:
kwargs["ltx2_vae_tiling"] = normalized.pipeline.vae_tiling
if engine.compile.text_encoder_enabled is not None:
# ``FastVideoArgs.from_kwargs`` filters to declared fields, so
# this is a no-op on the current legacy path. Emit anyway so the
# realtime runtime (PR 7.6) — which reads from the kwargs dict
# before FastVideoArgs filtering — can pick it up once wired.
kwargs["enable_torch_compile_text_encoder"] = (engine.compile.text_encoder_enabled)
quantization = engine.quantization
if quantization is not None and quantization.text_encoder_quant is not None:
@@ -220,8 +265,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
kwargs["init_weights_from_safetensors"] = components.transformer_weights
if components.transformer_2_weights is not None:
kwargs["init_weights_from_safetensors_2"] = components.transformer_2_weights
if components.config_root is not None:
kwargs["config_model_path"] = components.config_root
if components.upsampler_weights is not None:
kwargs["ltx2_refine_upsampler_path"] = components.upsampler_weights
kwargs.update(deepcopy(normalized.pipeline.preset_overrides))
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
refine = preset_overrides.pop("refine", None)
if isinstance(refine, Mapping):
for key in _LTX2_REFINE_FLAT_KEYS:
if key in refine:
kwargs[f"ltx2_refine_{key}"] = refine[key]
kwargs.update(preset_overrides)
kwargs.update(deepcopy(normalized.pipeline.experimental))
return FastVideoArgs.from_kwargs(**kwargs)
@@ -271,10 +326,13 @@ def request_to_sampling_param(
) -> SamplingParam:
if request.plan is not None:
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
if request.state is not None:
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
sampling_param = SamplingParam.from_pretrained(model_path)
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():
@@ -316,6 +374,25 @@ 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 {}
@@ -476,6 +553,37 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
_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,
@@ -509,6 +617,7 @@ __all__ = [
"load_generator_config_from_file",
"normalize_generation_request",
"normalize_generator_config",
"register_continuation_kind",
"request_to_pipeline_overrides",
"request_to_sampling_param",
]
+21 -6
View File
@@ -1,11 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import copy
from dataclasses import dataclass, field, fields
from typing import Any
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__)
@@ -92,9 +97,13 @@ class SamplingParam:
movement_distance: float | None = None
camera_rotation: str | None = None
# LTX2 multi-modal CFG and STG
ltx2_cfg_scale_video: float = 3.0
ltx2_cfg_scale_audio: float = 7.0
# 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
@@ -103,6 +112,12 @@ class SamplingParam:
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
@@ -127,7 +142,7 @@ class SamplingParam:
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str) -> "SamplingParam":
def from_pretrained(cls, model_path: str) -> SamplingParam:
sampling_param = cls._from_preset(model_path)
if sampling_param is not None:
return sampling_param
@@ -143,7 +158,7 @@ class SamplingParam:
def _from_preset(
cls,
model_path: str,
) -> "SamplingParam | None":
) -> SamplingParam | None:
"""Build a SamplingParam from preset defaults.
Returns ``None`` when no preset is configured for
+20 -1
View File
@@ -33,8 +33,25 @@ class OffloadConfig:
@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
kwargs: dict[str, Any] = field(default_factory=dict)
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
@@ -76,6 +93,8 @@ class PipelineSelection:
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)
+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)
+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, ))
+29 -2
View File
@@ -1,4 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.entrypoints.streaming.server import run_server
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__ = ["run_server"]
__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",
]
+520 -4
View File
@@ -1,15 +1,531 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-generator FastAPI + WebSocket streaming server."""
from __future__ import annotations
from fastvideo.api.schema import ServeConfig
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
def run_server(serve_config: ServeConfig) -> None:
"""Launch the streaming (WebSocket / Dynamo) server."""
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.")
raise NotImplementedError("streaming server is not implemented yet")
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",
]
+214
View File
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-connection session lifecycle for the streaming server.
Each WebSocket opens exactly one :class:`Session`. :class:`SessionManager`
enforces the ``generation_segment_cap`` and ``session_timeout_seconds``
budgets from :class:`fastvideo.api.StreamingConfig`.
"""
from __future__ import annotations
import enum
import time
import uuid
from dataclasses import dataclass, field
from typing import Any
from fastvideo.api.schema import ContinuationState
class SessionState(enum.Enum):
"""State-machine positions for a streaming session.
Transitions are server-owned. See
``docs/design/server_contracts/streaming.md`` for the full diagram.
"""
INITIALIZING = "initializing"
QUEUED = "queued"
GPU_BINDING = "gpu_binding"
ACTIVE = "active"
COMPLETE = "complete"
ERROR = "error"
TIMEOUT = "timeout"
REJECTED = "rejected"
_VALID_TRANSITIONS: dict[SessionState, frozenset[SessionState]] = {
SessionState.INITIALIZING:
frozenset({
SessionState.QUEUED,
SessionState.GPU_BINDING,
SessionState.REJECTED,
SessionState.ERROR,
}),
SessionState.QUEUED:
frozenset({
SessionState.GPU_BINDING,
SessionState.ERROR,
SessionState.TIMEOUT,
SessionState.REJECTED,
}),
SessionState.GPU_BINDING:
frozenset({
SessionState.ACTIVE,
SessionState.ERROR,
SessionState.TIMEOUT,
}),
SessionState.ACTIVE:
frozenset({
SessionState.ACTIVE,
SessionState.COMPLETE,
SessionState.ERROR,
SessionState.TIMEOUT,
}),
SessionState.COMPLETE:
frozenset(),
SessionState.ERROR:
frozenset(),
SessionState.TIMEOUT:
frozenset(),
SessionState.REJECTED:
frozenset(),
}
class InvalidSessionTransition(RuntimeError):
"""Raised when a session is asked to transition along an illegal edge."""
@dataclass
class Session:
id: str = field(default_factory=lambda: uuid.uuid4().hex)
state: SessionState = SessionState.INITIALIZING
created_at: float = field(default_factory=time.monotonic)
last_activity: float = field(default_factory=time.monotonic)
client_id: str | None = None
preset: str | None = None
preset_label: str | None = None
curated_prompts: list[str] = field(default_factory=list)
segment_idx: int = 0
enhancement_enabled: bool = False
auto_extension_enabled: bool = False
loop_generation_enabled: bool = False
single_clip_mode: bool = False
generation_paused: bool = False
stream_mode: str = "av_fmp4"
gpu_id: int | None = None
continuation_state: ContinuationState | None = None
metadata: dict[str, Any] = field(default_factory=dict)
def transition(self, target: SessionState) -> None:
"""Move to ``target`` if the edge is allowed.
Raises :class:`InvalidSessionTransition` on illegal moves. The
self-loop on ``ACTIVE`` is legal so the server can re-assert
ACTIVE on segment completion without special casing.
"""
allowed = _VALID_TRANSITIONS.get(self.state, frozenset())
if target not in allowed and target is not self.state:
raise InvalidSessionTransition(f"{self.state.value} -> {target.value} is not a valid "
f"session transition")
self.state = target
self.last_activity = time.monotonic()
def touch(self) -> None:
self.last_activity = time.monotonic()
def is_active(self) -> bool:
return self.state is SessionState.ACTIVE
def segment_cap_reached(self, cap: int) -> bool:
return self.segment_idx >= cap
class SessionManager:
"""Registers sessions and enforces per-server session limits."""
def __init__(
self,
*,
segment_cap: int,
session_timeout_seconds: int,
max_sessions: int = 1,
) -> None:
self._segment_cap = segment_cap
self._session_timeout_seconds = session_timeout_seconds
self._max_sessions = max_sessions
self._sessions: dict[str, Session] = {}
@property
def segment_cap(self) -> int:
return self._segment_cap
@property
def session_timeout_seconds(self) -> int:
return self._session_timeout_seconds
def create(self) -> Session:
if len(self._sessions) >= self._max_sessions:
raise SessionRejected(f"max sessions reached ({self._max_sessions})")
session = Session()
self._sessions[session.id] = session
return session
def get(self, session_id: str) -> Session | None:
return self._sessions.get(session_id)
def close(self, session_id: str) -> None:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
return session_id in self._sessions
def __len__(self) -> int:
return len(self._sessions)
def active_sessions(self) -> list[Session]:
return [s for s in self._sessions.values() if s.is_active()]
def reap_timed_out(self, now: float | None = None) -> list[str]:
"""Return the ids of sessions that have exceeded the idle timeout.
The caller is responsible for actually closing them — this
method only *identifies* dead sessions so the server can emit
``session_timeout`` frames before dropping the WebSocket.
TODO: unused until a background driver calls it. Per-connection
idle enforcement currently happens via asyncio.wait_for on
receive_json; this helper catches sessions stuck before any
receive (e.g. future QUEUED state) and is expected to be wired
into the GPU-pool reaper.
"""
now = now if now is not None else time.monotonic()
dead: list[str] = []
for sid, session in self._sessions.items():
if session.state in {
SessionState.COMPLETE,
SessionState.ERROR,
SessionState.TIMEOUT,
SessionState.REJECTED,
}:
continue
if now - session.last_activity > self._session_timeout_seconds:
dead.append(sid)
return dead
class SessionRejected(RuntimeError):
"""Raised when session creation fails (queue full, auth, etc.)."""
__all__ = [
"InvalidSessionTransition",
"Session",
"SessionManager",
"SessionRejected",
"SessionState",
]
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
"""Persist the initial-image blob attached to a streaming session."""
from __future__ import annotations
import base64
import binascii
import contextlib
import os
import tempfile
from dataclasses import dataclass
from typing import Any
_ACCEPTED_MIMES = {
"image/png": ".png",
"image/jpeg": ".jpg",
"image/jpg": ".jpg",
"image/webp": ".webp",
}
_MAX_IMAGE_BYTES = 32 * 1024 * 1024 # 32 MiB cap
@dataclass(frozen=True)
class SessionInitImage:
"""Location of the persisted init image.
Callers pass ``path`` to ``InputConfig.image_path``; ``display_name``
is only used for logs.
"""
path: str
display_name: str
mime: str
def persist_session_init_image(
payload: Any,
*,
output_dir: str | None = None,
) -> SessionInitImage | None:
"""Decode a client init-image blob and persist it to disk.
``payload`` shape (matches the internal UI protocol)::
{
"mime": "image/png",
"name": "ref.png",
"data": "<base64 bytes>",
}
Returns ``None`` when ``payload`` is falsy (no init image). Raises
:class:`ValueError` on schema / size / decode errors so the caller
can surface a user-facing ``error`` frame.
"""
if not payload:
return None
if not isinstance(payload, dict):
raise ValueError("session init image must be an object")
mime = payload.get("mime")
if mime not in _ACCEPTED_MIMES:
raise ValueError(f"session init image mime {mime!r} is not one of "
f"{sorted(_ACCEPTED_MIMES)}")
data_b64 = payload.get("data")
if not isinstance(data_b64, str):
raise ValueError("session init image data must be a base64 string")
try:
data = base64.b64decode(data_b64, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"session init image data is not valid base64: {exc}") from exc
if len(data) > _MAX_IMAGE_BYTES:
raise ValueError(f"session init image is {len(data)} bytes; limit is "
f"{_MAX_IMAGE_BYTES}")
if len(data) == 0:
raise ValueError("session init image data is empty")
ext = _ACCEPTED_MIMES[mime]
display_name = _sanitize_display_name(payload.get("name")) or f"init{ext}"
fd, path = tempfile.mkstemp(prefix="fastvideo-init-", suffix=ext, dir=output_dir)
try:
with os.fdopen(fd, "wb") as f:
f.write(data)
except Exception:
with contextlib.suppress(FileNotFoundError):
os.unlink(path)
raise
return SessionInitImage(path=path, display_name=display_name, mime=mime)
def _sanitize_display_name(name: Any) -> str | None:
if not isinstance(name, str):
return None
name = name.strip()
if not name:
return None
# Strip any path components — we only keep the leaf for logging.
return os.path.basename(name)
__all__ = [
"SessionInitImage",
"persist_session_init_image",
]
@@ -0,0 +1,206 @@
# SPDX-License-Identifier: Apache-2.0
"""Session state store for the FastVideo streaming server.
The streaming server keeps continuation state (decoded frames + audio
latents from the previous segment) server-side so the client doesn't
re-upload multi-megabyte tensors each WebSocket message. Two operations
are needed:
* ``snapshot(session_id) -> ContinuationState`` — serialize the current
state so it can be exported (e.g. over HTTP) or migrated to a
different server.
* ``hydrate(state) -> session_id`` — load a previously serialized state
into a new session (for resume-after-disconnect flows).
The store is an ABC with an :class:`InMemorySessionStore` default; Redis
or other backends can drop in without touching the pipeline.
Large tensor payloads (video frames, audio latents) are kept out of the
JSON payload via an accompanying :class:`BlobStore`. Both stores share a
process today; they are separate types so that a future implementation
can put blobs on S3 while keeping session metadata in Redis.
"""
from __future__ import annotations
import threading
import uuid
from abc import ABC, abstractmethod
from collections.abc import Iterator
from dataclasses import dataclass
from fastvideo.api.schema import ContinuationState
class BlobStore(ABC):
"""Opaque byte-blob storage keyed by id.
A :class:`ContinuationState` payload can reference large tensors
stored in a :class:`BlobStore` rather than inlining them, so the
JSON payload stays small when the state travels over the wire.
"""
@abstractmethod
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
"""Store ``data`` and return a blob id for later retrieval."""
@abstractmethod
def get(self, blob_id: str) -> bytes:
"""Load a previously stored blob. Raises ``KeyError`` if absent."""
@abstractmethod
def drop(self, blob_id: str) -> None:
"""Remove a blob. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, blob_id: str) -> bool:
...
@dataclass(frozen=True)
class _BlobRecord:
data: bytes
mime: str
class InMemoryBlobStore(BlobStore):
"""Thread-safe in-memory :class:`BlobStore` for single-process servers.
No eviction policy — callers are responsible for calling
:meth:`drop` when a blob's owning state is replaced or a session
ends. A redis- or filesystem-backed :class:`BlobStore` should
replace this when the streaming server lands as a real service
(PR 7.5+).
"""
def __init__(self) -> None:
self._blobs: dict[str, _BlobRecord] = {}
self._lock = threading.Lock()
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
blob_id = uuid.uuid4().hex
with self._lock:
self._blobs[blob_id] = _BlobRecord(data=data, mime=mime)
return blob_id
def get(self, blob_id: str) -> bytes:
with self._lock:
record = self._blobs.get(blob_id)
if record is None:
raise KeyError(f"Unknown blob id: {blob_id}")
return record.data
def drop(self, blob_id: str) -> None:
with self._lock:
self._blobs.pop(blob_id, None)
def __contains__(self, blob_id: str) -> bool:
with self._lock:
return blob_id in self._blobs
def __len__(self) -> int:
with self._lock:
return len(self._blobs)
class SessionStore(ABC):
"""Keyed store for per-session continuation state.
Implementations own the session-id → state mapping. The streaming
server calls :meth:`store` after each segment and :meth:`snapshot`
when a client explicitly asks for an exportable state handle.
"""
@abstractmethod
def store(self, session_id: str, state: ContinuationState) -> None:
"""Persist ``state`` for ``session_id``, replacing any prior value."""
@abstractmethod
def snapshot(self, session_id: str) -> ContinuationState | None:
"""Return the current state for ``session_id`` (or ``None``)."""
@abstractmethod
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
"""Install ``state`` as the starting point for a session.
When ``session_id`` is ``None`` the store allocates a fresh id
(UUID4); when provided the store uses it verbatim, overwriting
any prior state at that id.
"""
@abstractmethod
def drop(self, session_id: str) -> None:
"""Forget a session. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, session_id: str) -> bool:
...
@abstractmethod
def __iter__(self) -> Iterator[str]:
...
class InMemorySessionStore(SessionStore):
"""Thread-safe in-memory :class:`SessionStore`.
Default implementation used by single-process deployments; a future
Redis-backed store can be dropped in without changes to the server.
No eviction / TTL / bounded capacity — sessions only leave via
:meth:`drop`. The live streaming server (PR 7.5+) is responsible
for bounding growth and for dropping any :class:`BlobStore` blobs
referenced by a state when that state is replaced or a session
ends; this class does not know about blobs.
"""
def __init__(self) -> None:
self._sessions: dict[str, ContinuationState] = {}
self._lock = threading.Lock()
def store(self, session_id: str, state: ContinuationState) -> None:
with self._lock:
self._sessions[session_id] = state
def snapshot(self, session_id: str) -> ContinuationState | None:
with self._lock:
return self._sessions.get(session_id)
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
sid = session_id or uuid.uuid4().hex
with self._lock:
self._sessions[sid] = state
return sid
def drop(self, session_id: str) -> None:
with self._lock:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
with self._lock:
return session_id in self._sessions
def __iter__(self) -> Iterator[str]:
with self._lock:
return iter(list(self._sessions))
def __len__(self) -> int:
with self._lock:
return len(self._sessions)
__all__ = [
"BlobStore",
"InMemoryBlobStore",
"InMemorySessionStore",
"SessionStore",
]
+213
View File
@@ -0,0 +1,213 @@
# SPDX-License-Identifier: Apache-2.0
"""fMP4 stream encoder used by the streaming server.
The client's Media Source Extensions player needs a continuous fMP4
byte stream: first an *initialization segment* (``ftyp`` + ``moov``),
then one or more *media segments* (``moof`` + ``mdat``). We pipe raw
RGB frames into an ffmpeg subprocess configured for fragmented output
via ``-movflags empty_moov+default_base_moof+frag_keyframe+faststart``
and stream the bytes back out.
"""
from __future__ import annotations
import asyncio
import contextlib
import subprocess
import uuid
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
if TYPE_CHECKING:
import numpy as np
@dataclass
class FragmentedMP4Chunk:
"""A single fMP4 byte chunk emitted by :class:`FragmentedMP4Encoder`.
``kind`` identifies whether the chunk is the init segment (must be
fed into the client's ``SourceBuffer`` first) or a media fragment.
"""
kind: Literal["init", "media"]
data: bytes
stream_id: str
segment_idx: int
class FragmentedMP4Encoder:
"""Stream RGB frames in, fMP4 chunks out.
One encoder covers one segment. The server creates a new encoder
per :class:`ltx2_segment_start`` boundary so each segment becomes
one media fragment the client can append independently.
Example::
encoder = FragmentedMP4Encoder(width=1024, height=576, fps=24,
segment_idx=0)
async with encoder:
async for chunk in encoder.encode(frames):
await websocket.send_bytes(chunk.data)
"""
def __init__(
self,
*,
width: int,
height: int,
fps: int,
segment_idx: int,
stream_id: str | None = None,
ffmpeg_path: str = "ffmpeg",
preset: str = "ultrafast",
pixel_format_out: str = "yuv420p",
extra_args: list[str] | None = None,
) -> None:
self.width = width
self.height = height
self.fps = fps
self.segment_idx = segment_idx
self.stream_id = stream_id or uuid.uuid4().hex
self._ffmpeg_path = ffmpeg_path
self._preset = preset
self._pixel_format_out = pixel_format_out
self._extra_args = list(extra_args or [])
self._proc: subprocess.Popen | None = None
self._init_emitted = False
async def __aenter__(self) -> FragmentedMP4Encoder:
self._spawn()
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
await self.close()
def _spawn(self) -> None:
args = [
self._ffmpeg_path,
"-hide_banner",
"-loglevel",
"error",
"-f",
"rawvideo",
"-pix_fmt",
"rgb24",
"-s",
f"{self.width}x{self.height}",
"-r",
str(self.fps),
"-i",
"-",
"-c:v",
"libx264",
"-preset",
self._preset,
"-tune",
"zerolatency",
"-pix_fmt",
self._pixel_format_out,
"-movflags",
"empty_moov+default_base_moof+frag_keyframe+faststart",
"-f",
"mp4",
*self._extra_args,
"-",
]
# stderr → DEVNULL: with -loglevel error on, the only thing
# stderr would carry is unsolicited warnings. Piping without a
# reader deadlocks ffmpeg once the pipe buffer (~64 KiB) fills.
self._proc = subprocess.Popen( # noqa: S603
args,
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL,
bufsize=0,
)
async def encode(
self,
frames: list[np.ndarray] | AsyncIterator[np.ndarray],
) -> AsyncIterator[FragmentedMP4Chunk]:
"""Feed frames into ffmpeg and yield fMP4 chunks as they appear."""
if self._proc is None:
self._spawn()
assert self._proc is not None and self._proc.stdin is not None
proc = self._proc
loop = asyncio.get_running_loop()
async def _writer() -> None:
try:
if hasattr(frames, "__aiter__"):
async for frame in frames: # type: ignore[union-attr]
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
else:
for frame in frames: # type: ignore[assignment]
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
finally:
with contextlib.suppress(BrokenPipeError):
proc.stdin.close()
writer_task = asyncio.create_task(_writer())
try:
reader = proc.stdout
assert reader is not None
# Read in reasonably-sized chunks; MSE tolerates any size
# but we don't want to starve the event loop.
chunk_size = 64 * 1024
while True:
data = await loop.run_in_executor(None, reader.read, chunk_size)
if not data:
break
kind: Literal["init", "media"] = "init" if not self._init_emitted else "media"
self._init_emitted = True
yield FragmentedMP4Chunk(
kind=kind,
data=bytes(data),
stream_id=self.stream_id,
segment_idx=self.segment_idx,
)
finally:
await writer_task
async def close(self) -> None:
if self._proc is None:
return
proc = self._proc
self._proc = None
try:
if proc.stdin and not proc.stdin.closed:
proc.stdin.close()
except BrokenPipeError:
pass
loop = asyncio.get_running_loop()
try:
await asyncio.wait_for(
loop.run_in_executor(None, proc.wait),
timeout=5.0,
)
except asyncio.TimeoutError:
proc.kill()
await loop.run_in_executor(None, proc.wait)
def _write_frame(stdin, frame: np.ndarray) -> None:
import numpy as np
if not isinstance(frame, np.ndarray):
raise TypeError("fMP4 encoder frames must be numpy.ndarray")
if frame.dtype != np.uint8:
frame = frame.astype(np.uint8)
if frame.ndim != 3 or frame.shape[-1] != 3:
raise ValueError("fMP4 encoder frames must be HxWx3 uint8 RGB; got "
f"shape={frame.shape}, dtype={frame.dtype}")
with contextlib.suppress(BrokenPipeError):
stdin.write(frame.tobytes())
__all__ = [
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
]
+11 -1
View File
@@ -844,6 +844,13 @@ class TrainingArgs(FastVideoArgs):
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
# CFG scale applied to the real (teacher) score in the DMD loss, using the
# parameterization `x = x_cond + w * (x_cond - x_uncond)`. This differs
# from the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)` by an
# offset of 1: `w_here = w_standard - 1`. So `w=0` recovers the
# conditional output, `w=-1` recovers the unconditional output, and the
# default 3.5 corresponds to a standard CFG scale of 4.5. Matches the
# original DMD2 reference implementation.
real_score_guidance_scale: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
@@ -1104,7 +1111,10 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--real-score-guidance-scale",
type=float,
default=TrainingArgs.real_score_guidance_scale,
help="Teacher guidance scale")
help=("Teacher CFG scale for the real score in the DMD loss. Uses "
"the parameterization x_cond + w * (x_cond - x_uncond), so "
"w=0 -> cond, w=-1 -> uncond, and the relation to standard "
"CFG is w_standard = w + 1 (default 3.5 == standard 4.5)."))
parser.add_argument("--fake-score-learning-rate",
type=float,
default=TrainingArgs.fake_score_learning_rate,
@@ -0,0 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# Importing continuation registers the "ltx2.v1" continuation kind with
# the public compat layer so GenerationRequest.state(kind="ltx2.v1") is
# recognized on the public API boundary.
from fastvideo.pipelines.basic.ltx2 import continuation # noqa: F401
@@ -0,0 +1,386 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed continuation state for the LTX-2 streaming pipeline.
Segment N+1 conditions on segment N's trailing decoded frames and
denoised audio latents. The streaming runtime used to hold this state as
per-worker globals; lifting it into a typed, JSON-serializable object
lets clients snapshot, migrate, or round-trip it through an HTTP/RPC
boundary. The envelope ``ContinuationState(kind, payload)`` is the
shared public API; the typed class here owns the LTX-2 payload shape.
Serialization contract:
* Video frames → PNG bytes + base64, or a :class:`BlobStore` id.
* Audio latents → a self-describing safetensors blob + base64, or a
:class:`BlobStore` id. safetensors preserves ``bfloat16``, which a
raw-numpy round-trip cannot.
* The returned payload is always a plain JSON-serializable dict.
"""
from __future__ import annotations
import base64
from collections.abc import Mapping
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from fastvideo.api.compat import register_continuation_kind
from fastvideo.api.schema import ContinuationState
if TYPE_CHECKING:
import numpy as np
import torch
from fastvideo.entrypoints.streaming.session_store import BlobStore
LTX2_CONTINUATION_KIND = "ltx2.v1"
"""Public ``ContinuationState.kind`` for LTX-2 payloads."""
LTX2_CONTINUATION_SCHEMA_VERSION = 1
"""Payload schema version carried inside ``payload.schema_version``."""
DEFAULT_INLINE_THRESHOLD_BYTES = 2 * 1024 * 1024
"""Tensors larger than this go to the blob store (if available). 2 MiB
is below typical single-JSON-message limits (Dynamo: 4 MiB, Postgres
TOAST: 1 GiB) and well above per-frame PNG payloads (~200 KiB at
512x512)."""
@dataclass
class LTX2ContinuationState:
"""Typed LTX-2 continuation state carried between streaming segments.
``video_frames`` hold trailing decoded RGB frames (uint8 HxWx3) from
segment N for conditioning segment N+1 via the VAE encode path.
``audio_latents`` is the cached denoised audio latent tensor of shape
``[B, C, T, mel]`` that segment N+1 will copy into the overlap
region of its clean-latent conditioning.
Most fields map 1:1 onto the internal gpu_pool's per-worker state;
the only new concept is the ``*_blob_id`` fields, which allow large
tensors to live outside the JSON payload. See module docstring.
"""
segment_index: int = 0
"""Index of the *just-completed* segment. Segment 0 has no history;
state returned after segment 0 carries ``segment_index=0`` and the
caller uses ``segment_index + 1`` as the next segment number."""
video_frames: list[np.ndarray] | None = None
"""Trailing decoded frames, each an RGB uint8 ``np.ndarray`` shaped
``(H, W, 3)``. ``None`` when the state is blob-backed or unset."""
video_frames_blob_id: str | None = None
"""Blob store id when the frames live outside the payload."""
video_conditioning_frame_idx: int = 0
"""Target frame index inside the next segment that the trailing
frames align with (matches the LTX-2 ``ltx2_video_conditions``
tuple's ``frame_idx`` slot)."""
video_conditioning_strength: float = 1.0
"""Conditioning strength in [0, 1]. Matches the ``ltx2_video_
conditions`` tuple's strength slot."""
audio_latents: torch.Tensor | None = None
"""Denoised audio latent tensor of shape ``[B, C, T, mel]``.
``None`` when the state is blob-backed or unset."""
audio_latents_blob_id: str | None = None
"""Blob store id when audio latents live outside the payload."""
audio_sample_rate: int | None = None
"""Sample rate for the audio side (e.g. 24000)."""
audio_conditioning_num_frames: int = 0
"""Number of trailing audio frames that carry over as clean context
into segment N+1."""
audio_conditioning_strength: float = 1.0
"""Clean-latent mask value applied to the overlap region; 0.0 keeps
the cached audio entirely, 1.0 renoises from scratch."""
video_position_offset_sec: float = 0.0
"""Seconds by which video RoPE is shifted forward so the audio
prefix can sit at ``t >= 0`` when audio conditioning is longer than
video conditioning."""
metadata: dict[str, Any] = field(default_factory=dict)
"""Opaque metadata bag for forward-compat fields that don't need
their own typed slot yet (e.g. custom knob experiments)."""
def to_continuation_state(
self,
*,
blob_store: BlobStore | None = None,
inline_threshold_bytes: int = DEFAULT_INLINE_THRESHOLD_BYTES,
) -> ContinuationState:
"""Serialize into a public :class:`ContinuationState`.
When ``blob_store`` is given, tensors larger than
``inline_threshold_bytes`` are stored via
:meth:`BlobStore.put` and referenced by id; otherwise all data
is base64-encoded inline. The payload is always a plain
JSON-serializable dict.
"""
payload: dict[str, Any] = {
"schema_version": LTX2_CONTINUATION_SCHEMA_VERSION,
"segment_index": int(self.segment_index),
"video_conditioning_frame_idx": int(self.video_conditioning_frame_idx),
"video_conditioning_strength": float(self.video_conditioning_strength),
"audio_conditioning_num_frames": int(self.audio_conditioning_num_frames),
"audio_conditioning_strength": float(self.audio_conditioning_strength),
"video_position_offset_sec": float(self.video_position_offset_sec),
"metadata": dict(self.metadata),
}
if self.audio_sample_rate is not None:
payload["audio_sample_rate"] = int(self.audio_sample_rate)
video_payload = self._encode_video_frames(
blob_store=blob_store,
inline_threshold_bytes=inline_threshold_bytes,
)
if video_payload is not None:
payload["video"] = video_payload
audio_payload = self._encode_audio_latents(
blob_store=blob_store,
inline_threshold_bytes=inline_threshold_bytes,
)
if audio_payload is not None:
payload["audio"] = audio_payload
return ContinuationState(
kind=LTX2_CONTINUATION_KIND,
payload=payload,
)
@classmethod
def from_continuation_state(
cls,
state: ContinuationState,
*,
blob_store: BlobStore | None = None,
) -> LTX2ContinuationState:
"""Rebuild a typed state from a public :class:`ContinuationState`.
Raises :class:`ValueError` when the kind doesn't match or the
schema version is unsupported.
"""
if state.kind != LTX2_CONTINUATION_KIND:
raise ValueError(f"Expected ContinuationState.kind={LTX2_CONTINUATION_KIND!r}, "
f"got {state.kind!r}")
payload = state.payload or {}
version = int(payload.get("schema_version", LTX2_CONTINUATION_SCHEMA_VERSION))
if version != LTX2_CONTINUATION_SCHEMA_VERSION:
raise ValueError(f"Unsupported LTX-2 continuation schema_version={version}; "
f"this build expects {LTX2_CONTINUATION_SCHEMA_VERSION}")
out = cls(
segment_index=int(payload.get("segment_index", 0)),
video_conditioning_frame_idx=int(payload.get("video_conditioning_frame_idx", 0)),
video_conditioning_strength=float(payload.get("video_conditioning_strength", 1.0)),
audio_sample_rate=(int(payload["audio_sample_rate"]) if "audio_sample_rate" in payload else None),
audio_conditioning_num_frames=int(payload.get("audio_conditioning_num_frames", 0)),
audio_conditioning_strength=float(payload.get("audio_conditioning_strength", 1.0)),
video_position_offset_sec=float(payload.get("video_position_offset_sec", 0.0)),
metadata=dict(payload.get("metadata") or {}),
)
video = payload.get("video")
if isinstance(video, Mapping):
cls._decode_video_frames(out, video, blob_store=blob_store)
audio = payload.get("audio")
if isinstance(audio, Mapping):
cls._decode_audio_latents(out, audio, blob_store=blob_store)
return out
# ------------------------------------------------------------------
# Video frame helpers
# ------------------------------------------------------------------
def _encode_video_frames(
self,
*,
blob_store: BlobStore | None,
inline_threshold_bytes: int,
) -> dict[str, Any] | None:
if self.video_frames_blob_id is not None:
return {"blob_id": self.video_frames_blob_id}
if not self.video_frames:
return None
encoded = [_encode_png(frame) for frame in self.video_frames]
total = sum(len(b) for b in encoded)
if blob_store is not None and total > inline_threshold_bytes:
concatenated = _pack_frame_blobs(encoded)
blob_id = blob_store.put(
concatenated,
mime="application/x-fastvideo-frames+png",
)
return {"blob_id": blob_id, "frame_count": len(encoded)}
return {
"frames_b64": [base64.b64encode(b).decode("ascii") for b in encoded],
}
@staticmethod
def _decode_video_frames(
out: LTX2ContinuationState,
video: Mapping[str, Any],
*,
blob_store: BlobStore | None,
) -> None:
blob_id = video.get("blob_id")
if isinstance(blob_id, str):
if blob_store is None:
out.video_frames_blob_id = blob_id
return
raw = blob_store.get(blob_id)
encoded = _unpack_frame_blobs(raw)
out.video_frames = [_decode_png(b) for b in encoded]
return
frames_b64 = video.get("frames_b64")
if isinstance(frames_b64, list):
decoded = [_decode_png(base64.b64decode(b)) for b in frames_b64 if isinstance(b, str)]
out.video_frames = decoded or None
# ------------------------------------------------------------------
# Audio latent helpers
# ------------------------------------------------------------------
def _encode_audio_latents(
self,
*,
blob_store: BlobStore | None,
inline_threshold_bytes: int,
) -> dict[str, Any] | None:
if self.audio_latents_blob_id is not None:
return {"blob_id": self.audio_latents_blob_id}
if self.audio_latents is None:
return None
raw = _tensor_to_safetensors_bytes(self.audio_latents)
if blob_store is not None and len(raw) > inline_threshold_bytes:
blob_id = blob_store.put(
raw,
mime="application/x-fastvideo-tensor+safetensors",
)
return {"blob_id": blob_id}
return {"safetensors_b64": base64.b64encode(raw).decode("ascii")}
@staticmethod
def _decode_audio_latents(
out: LTX2ContinuationState,
audio: Mapping[str, Any],
*,
blob_store: BlobStore | None,
) -> None:
blob_id = audio.get("blob_id")
if isinstance(blob_id, str):
if blob_store is None:
out.audio_latents_blob_id = blob_id
return
raw = blob_store.get(blob_id)
out.audio_latents = _safetensors_bytes_to_tensor(raw)
return
data_b64 = audio.get("safetensors_b64")
if isinstance(data_b64, str):
out.audio_latents = _safetensors_bytes_to_tensor(base64.b64decode(data_b64))
def _encode_png(frame: np.ndarray) -> bytes:
"""Encode an ``(H, W, 3)`` uint8 RGB frame as PNG bytes."""
import numpy as np
from PIL import Image
if not isinstance(frame, np.ndarray):
raise TypeError(f"LTX2 continuation frame must be a numpy ndarray, got {type(frame).__name__}")
if frame.dtype != np.uint8 or frame.ndim != 3 or frame.shape[-1] != 3:
raise ValueError("LTX2 continuation frame must be uint8 HxWx3 RGB; got "
f"dtype={frame.dtype}, shape={frame.shape}")
import io
buffer = io.BytesIO()
Image.fromarray(frame).save(buffer, format="PNG")
return buffer.getvalue()
def _decode_png(data: bytes) -> np.ndarray:
import io
import numpy as np
from PIL import Image
img = Image.open(io.BytesIO(data)).convert("RGB")
return np.array(img, dtype=np.uint8)
def _pack_frame_blobs(encoded: list[bytes]) -> bytes:
"""Pack multiple PNG blobs into a single blob for blob-store storage.
Format: ``[4-byte big-endian count][4-byte len][png][4-byte len][png]...``.
"""
parts: list[bytes] = [len(encoded).to_bytes(4, "big")]
for blob in encoded:
parts.append(len(blob).to_bytes(4, "big"))
parts.append(blob)
return b"".join(parts)
def _unpack_frame_blobs(raw: bytes) -> list[bytes]:
if len(raw) < 4:
raise ValueError("frame blob truncated: missing count header")
count = int.from_bytes(raw[:4], "big")
# Each frame contributes at least a 4-byte length prefix, so a
# declared count larger than (len(raw) - 4) // 4 cannot fit and
# would otherwise cause an O(count) allocation loop on malformed
# input.
if count > (len(raw) - 4) // 4:
raise ValueError(f"frame blob declares {count} frames but buffer holds at most "
f"{(len(raw) - 4) // 4}")
out: list[bytes] = []
cursor = 4
for index in range(count):
if cursor + 4 > len(raw):
raise ValueError(f"frame blob truncated at frame {index} length header")
length = int.from_bytes(raw[cursor:cursor + 4], "big")
cursor += 4
if cursor + length > len(raw):
raise ValueError(f"frame blob truncated at frame {index} payload")
out.append(raw[cursor:cursor + length])
cursor += length
return out
def _tensor_to_safetensors_bytes(tensor: Any) -> bytes:
"""Serialize a torch tensor to a self-describing safetensors blob.
Uses the in-memory safetensors API so the wire format preserves
dtype (including ``bfloat16``, which a raw-numpy path cannot) and
shape without needing sidecar metadata.
"""
import torch
from safetensors.torch import save as st_save
if isinstance(tensor, torch.Tensor):
return st_save({"t": tensor.detach().cpu()})
import numpy as np
if isinstance(tensor, np.ndarray):
return st_save({"t": torch.from_numpy(np.ascontiguousarray(tensor))})
raise TypeError("LTX2 audio_latents must be a torch.Tensor or numpy.ndarray, got "
f"{type(tensor).__name__}")
def _safetensors_bytes_to_tensor(raw: bytes) -> Any:
from safetensors.torch import load as st_load
return st_load(raw)["t"]
register_continuation_kind(LTX2_CONTINUATION_KIND)
__all__ = [
"DEFAULT_INLINE_THRESHOLD_BYTES",
"LTX2ContinuationState",
"LTX2_CONTINUATION_KIND",
"LTX2_CONTINUATION_SCHEMA_VERSION",
]
+35 -1
View File
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX2 model family pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_stage_override_fields, )
_LTX2_NEGATIVE_PROMPT = ("blurry, out of focus, overexposed, underexposed, low contrast, "
"washed out colors, excessive noise, grainy texture, poor lighting, "
@@ -30,6 +32,13 @@ _DENOISE_STAGE = PresetStageSpec(
}),
)
_REFINE_STAGE = PresetStageSpec(
name="refine",
kind="refinement",
description="Latent-upsample + second-pass refine",
allowed_overrides=refine_stage_override_fields(),
)
LTX2_BASE = InferencePreset(
name="ltx2_base",
version=1,
@@ -77,4 +86,29 @@ LTX2_DISTILLED = InferencePreset(
},
)
ALL_PRESETS = (LTX2_BASE, LTX2_DISTILLED)
LTX2_TWO_STAGE = InferencePreset(
name="ltx2_two_stage",
version=1,
model_family="ltx2",
description="LTX-2 distilled with 2x spatial refine (stage 1 half-res + stage 2 upsample+denoise)",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, _REFINE_STAGE),
defaults={
"seed": 10,
"height": 1024,
"width": 1536,
"num_frames": 121,
"fps": 24,
"guidance_scale": 1.0,
"num_inference_steps": 8,
"negative_prompt": "",
},
stage_defaults={
"refine": {
"num_inference_steps": 2,
"guidance_scale": 1.0,
},
},
)
ALL_PRESETS = (LTX2_BASE, LTX2_DISTILLED, LTX2_TWO_STAGE)
@@ -0,0 +1,59 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed override surfaces for the LTX-2 two-stage refine flow.
* ``preset_overrides.refine`` — init-time knobs (see
:class:`LTX2RefinePresetOverride`).
* ``stage_overrides.refine`` — per-request knobs (see
:class:`LTX2RefineStageOverride`).
Asset paths live on :class:`~fastvideo.api.schema.ComponentConfig`
(``upsampler_weights`` and ``lora_path``).
"""
from __future__ import annotations
from dataclasses import asdict, dataclass, fields
from typing import Any
@dataclass
class LTX2RefinePresetOverride:
"""Init-time refine wiring under ``preset_overrides.refine``."""
enabled: bool | None = None
add_noise: bool | None = None
@dataclass
class LTX2RefineStageOverride:
"""Per-request refine tuning under ``stage_overrides.refine``."""
# Stage-2 refine only validates 2 (reduced) and 3 (official distilled)
# sigma schedules; other values raise at pipeline construction.
num_inference_steps: int | None = None
guidance_scale: float | None = None
image_crf: int | None = None
video_position_offset_sec: float | None = None
def refine_override_to_dict(override: LTX2RefinePresetOverride | LTX2RefineStageOverride, ) -> dict[str, Any]:
"""Serialise a refine override, dropping ``None`` entries so only
user-set fields reach ``preset_overrides.refine`` or
``stage_overrides.refine``."""
return {k: v for k, v in asdict(override).items() if v is not None}
def refine_preset_override_fields() -> frozenset[str]:
return frozenset(f.name for f in fields(LTX2RefinePresetOverride))
def refine_stage_override_fields() -> frozenset[str]:
return frozenset(f.name for f in fields(LTX2RefineStageOverride))
__all__ = [
"LTX2RefinePresetOverride",
"LTX2RefineStageOverride",
"refine_override_to_dict",
"refine_preset_override_fields",
"refine_stage_override_fields",
]
@@ -0,0 +1,17 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2 family pipeline stages."""
from fastvideo.pipelines.basic.ltx2.stages.ltx2_audio_decoding import (
LTX2AudioDecodingStage, )
from fastvideo.pipelines.basic.ltx2.stages.ltx2_denoising import (
LTX2DenoisingStage, )
from fastvideo.pipelines.basic.ltx2.stages.ltx2_latent_preparation import (
LTX2LatentPreparationStage, )
from fastvideo.pipelines.basic.ltx2.stages.ltx2_text_encoding import (
LTX2TextEncodingStage, )
__all__ = [
"LTX2AudioDecodingStage",
"LTX2DenoisingStage",
"LTX2LatentPreparationStage",
"LTX2TextEncodingStage",
]
@@ -26,7 +26,7 @@ MATRIXGAME_I2V = InferencePreset(
"fps": 25,
"guidance_scale": 1.0,
"num_inference_steps": 3,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -26,7 +26,7 @@ TURBO_T2V_1_3B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -44,7 +44,7 @@ TURBO_T2V_14B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -62,7 +62,7 @@ TURBO_I2V_A14B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -17,6 +17,8 @@ import torch
if TYPE_CHECKING:
from torchcodec.decoders import VideoDecoder
from fastvideo.api.schema import ContinuationState
import time
from collections import OrderedDict
@@ -206,6 +208,9 @@ class ForwardBatch:
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None
continuation_state: "ContinuationState | None" = None
return_continuation_state: bool = False
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
+6 -4
View File
@@ -25,10 +25,12 @@ from fastvideo.pipelines.stages.latent_preparation import (Cosmos25LatentPrepara
Cosmos25AutoLatentPreparationStage,
Cosmos25T2WLatentPreparationStage,
Cosmos25V2WLatentPreparationStage, LatentPreparationStage)
from fastvideo.pipelines.stages.ltx2_audio_decoding import LTX2AudioDecodingStage
from fastvideo.pipelines.stages.ltx2_denoising import LTX2DenoisingStage
from fastvideo.pipelines.stages.ltx2_latent_preparation import (LTX2LatentPreparationStage)
from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
from fastvideo.pipelines.basic.ltx2.stages import (
LTX2AudioDecodingStage,
LTX2DenoisingStage,
LTX2LatentPreparationStage,
LTX2TextEncodingStage,
)
from fastvideo.pipelines.stages.matrixgame_denoising import (MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
@@ -524,6 +524,22 @@ class ImageVAEEncodingStage(PipelineStage):
image = resize(image, height, width, resize_mode=resize_mode)
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
elif isinstance(image, torch.Tensor):
# VideoTransformStage delivers uint8 [0, 255] frames via batch.pil_image
# for the I2V preprocessing path. Convert here (not at the source) because
# batch.pil_image is also consumed as uint8 by ImageEncodingStage (HF
# processor does its own rescale) and by record_schema.py for parquet.
if image.dtype == torch.uint8:
image = image.float() / 255.0
elif not image.dtype.is_floating_point:
raise ValueError(f"preprocess() expected uint8 or float tensor, got {image.dtype}")
image_min = image.min()
image_max = image.max()
if image_max > 1.0 + 1e-4 or image_min < -1.0 - 1e-4:
raise ValueError("preprocess() expected tensor in [0, 1] or [-1, 1], got "
f"range [{image_min.item():.3f}, {image_max.item():.3f}]")
else:
raise TypeError(f"preprocess() expected PIL.Image or torch.Tensor, got {type(image)}")
do_normalize = True
if image.min() < 0:
+1 -1
View File
@@ -26,7 +26,7 @@ from fastvideo.configs.pipelines.hunyuan15 import (Hunyuan15T2V480PConfig, Hunyu
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionI2V_A14B_Config,
TurboDiffusionT2V_14B_Config,
+16 -4
View File
@@ -519,7 +519,7 @@ def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch)
cli_main.main()
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, monkeypatch):
config_path = tmp_path / "serve-streaming.yaml"
config_path.write_text(
"generator:\n"
@@ -530,9 +530,21 @@ def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
)
args, _ = _parse_serve_args(["--config", str(config_path)])
with pytest.raises(NotImplementedError,
match="streaming server is not implemented"):
ServeSubcommand().cmd(args)
captured: dict[str, object] = {}
def fake_run_server(serve_config, *, generator=None):
captured["serve_config"] = serve_config
def fail_if_called(*_args, **_kwargs):
raise AssertionError("OpenAI server must not run when streaming is set")
monkeypatch.setattr(streaming_server, "run_server", fake_run_server)
monkeypatch.setattr(api_server, "run_server", fail_if_called)
ServeSubcommand().cmd(args)
serve_config = captured["serve_config"]
assert serve_config.streaming is not None
assert serve_config.streaming.stream_mode == "av_fmp4"
def test_streaming_run_server_rejects_missing_streaming_block():
@@ -0,0 +1,226 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for ``fastvideo.api.compat`` translation helpers covering the
typed CompileConfig + PipelineSelection.vae_tiling surfaces promoted in
PR 6.
"""
from __future__ import annotations
from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
legacy_from_pretrained_to_config,
)
from fastvideo.api.schema import CompileConfig, GeneratorConfig
class TestLegacyTorchCompileKwargsTranslation:
"""Legacy ``torch_compile_kwargs={...}`` gets split across the four
first-class :class:`CompileConfig` fields and anything unknown falls
into ``extras``."""
def test_all_typed_keys_promoted(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{
"enable_torch_compile": True,
"torch_compile_kwargs": {
"backend": "inductor",
"fullgraph": True,
"mode": "max-autotune-no-cudagraphs",
"dynamic": False,
},
},
)
compile_config = config.engine.compile
assert compile_config.enabled is True
assert compile_config.backend == "inductor"
assert compile_config.fullgraph is True
assert compile_config.mode == "max-autotune-no-cudagraphs"
assert compile_config.dynamic is False
assert compile_config.extras == {}
def test_unknown_keys_land_in_extras(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{
"enable_torch_compile": True,
"torch_compile_kwargs": {
"backend": "inductor",
"options": {"triton.cudagraphs": False},
"disable": False,
},
},
)
compile_config = config.engine.compile
assert compile_config.backend == "inductor"
assert compile_config.extras == {
"options": {"triton.cudagraphs": False},
"disable": False,
}
def test_empty_kwargs_produces_empty_extras(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"torch_compile_kwargs": {}},
)
compile_config = config.engine.compile
assert compile_config.extras == {}
assert compile_config.backend is None
class TestCompileConfigRoundTrip:
"""typed CompileConfig -> FastVideoArgs.torch_compile_kwargs
reconstruction drops ``None`` typed fields and merges ``extras``."""
def test_only_typed_fields_emitted(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(
CompileConfig(enabled=True, backend="inductor", fullgraph=True)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["enable_torch_compile"] is True
assert args.kwargs["torch_compile_kwargs"] == {
"backend": "inductor",
"fullgraph": True,
}
def test_extras_merged_into_torch_compile_kwargs(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(
CompileConfig(
enabled=True,
mode="reduce-overhead",
extras={"options": {"triton.cudagraphs": False}},
)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["torch_compile_kwargs"] == {
"mode": "reduce-overhead",
"options": {"triton.cudagraphs": False},
}
def test_none_fields_suppressed(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["torch_compile_kwargs"] == {}
class TestLegacyLtx2VaeTilingTranslation:
"""``ltx2_vae_tiling`` flat kwarg promotes to
``generator.pipeline.vae_tiling``; reverse direction emits the
legacy name back to FastVideoArgs."""
def test_forward_routes_to_pipeline_vae_tiling(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"ltx2_vae_tiling": False},
)
assert config.pipeline.vae_tiling is False
def test_true_round_trips(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"ltx2_vae_tiling": True},
)
assert config.pipeline.vae_tiling is True
def test_unset_stays_none(self) -> None:
config = legacy_from_pretrained_to_config("/models/ltx2", {})
assert config.pipeline.vae_tiling is None
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
config.pipeline.vae_tiling = False
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["ltx2_vae_tiling"] is False
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
args = generator_config_to_fastvideo_args(config)
assert "ltx2_vae_tiling" not in args.kwargs
class TestLegacyTextEncoderCompileTranslation:
"""``enable_torch_compile_text_encoder`` flat kwarg promotes to
``generator.engine.compile.text_encoder_enabled``; reverse direction
emits the legacy name back onto the FastVideoArgs kwargs dict so
realtime-runtime consumers can read it before FastVideoArgs filters
unknown fields."""
def test_forward_routes_to_compile_text_encoder_enabled(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"enable_torch_compile_text_encoder": True},
)
assert config.engine.compile.text_encoder_enabled is True
def test_false_round_trips(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"enable_torch_compile_text_encoder": False},
)
assert config.engine.compile.text_encoder_enabled is False
def test_unset_stays_none(self) -> None:
config = legacy_from_pretrained_to_config("/models/ltx2", {})
assert config.engine.compile.text_encoder_enabled is None
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(
CompileConfig(text_encoder_enabled=True)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["enable_torch_compile_text_encoder"] is True
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
args = generator_config_to_fastvideo_args(config)
assert "enable_torch_compile_text_encoder" not in args.kwargs
# -------------------------------------------------------------------
# Helpers
# -------------------------------------------------------------------
def _engine_with_compile(compile_config):
"""Build an ``EngineConfig`` that carries the supplied compile block."""
from fastvideo.api.schema import EngineConfig
engine = EngineConfig()
engine.compile = compile_config
return engine
def _stub_fastvideo_args_from_kwargs(monkeypatch):
"""Swap ``FastVideoArgs.from_kwargs`` for a capture-only stub so
translation tests don't need to construct a valid FastVideoArgs."""
from fastvideo import fastvideo_args as fva
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _Captured)
@@ -0,0 +1,250 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the typed LTX-2 continuation state.
Covers:
* round-trip through :class:`ContinuationState` (inline and blob-backed)
* payload is JSON-serializable (Dynamo RPC / HTTP client constraint)
* kind / schema_version validation on deserialization
* compat-layer validation (known kinds, payload shape)
* round-trip through :func:`request_to_sampling_param` attaches the
state to the resulting :class:`SamplingParam` without losing fidelity
"""
from __future__ import annotations
import json
import numpy as np
import pytest
import torch
# Importing compat first, then the LTX-2 module, exercises the
# self-registration side effect on import (important for the API
# test suite where the pipeline package isn't otherwise imported).
from fastvideo.api import compat as api_compat # noqa: F401
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
OutputConfig,
)
from fastvideo.entrypoints.streaming.session_store import InMemoryBlobStore
from fastvideo.pipelines.basic.ltx2.continuation import (
LTX2_CONTINUATION_KIND,
LTX2_CONTINUATION_SCHEMA_VERSION,
LTX2ContinuationState,
)
def _make_typed_state() -> LTX2ContinuationState:
return LTX2ContinuationState(
segment_index=3,
video_frames=[
(np.ones((64, 64, 3), dtype=np.uint8) * (i * 10)) for i in range(4)
],
video_conditioning_frame_idx=9,
video_conditioning_strength=0.75,
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.float32),
audio_sample_rate=24000,
audio_conditioning_num_frames=5,
audio_conditioning_strength=0.5,
video_position_offset_sec=0.125,
metadata={"note": "unit-test"},
)
class TestRoundTrip:
"""Round-trip through :class:`ContinuationState` preserves all fields."""
def test_kind_and_schema_version(self):
state = _make_typed_state().to_continuation_state()
assert state.kind == LTX2_CONTINUATION_KIND
assert state.payload["schema_version"] == LTX2_CONTINUATION_SCHEMA_VERSION
def test_inline_roundtrip_preserves_scalars(self):
original = _make_typed_state()
envelope = original.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.segment_index == original.segment_index
assert restored.video_conditioning_frame_idx == (
original.video_conditioning_frame_idx)
assert restored.video_conditioning_strength == (
original.video_conditioning_strength)
assert restored.audio_sample_rate == original.audio_sample_rate
assert restored.audio_conditioning_num_frames == (
original.audio_conditioning_num_frames)
assert restored.audio_conditioning_strength == (
original.audio_conditioning_strength)
assert restored.video_position_offset_sec == (
original.video_position_offset_sec)
assert restored.metadata == original.metadata
def test_inline_roundtrip_preserves_video_frames(self):
original = _make_typed_state()
envelope = original.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.video_frames is not None
assert len(restored.video_frames) == len(original.video_frames)
for before, after in zip(original.video_frames,
restored.video_frames):
np.testing.assert_array_equal(before, after)
def test_inline_roundtrip_preserves_audio_latents(self):
original = _make_typed_state()
envelope = original.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.audio_latents is not None
assert tuple(restored.audio_latents.shape) == tuple(
original.audio_latents.shape)
assert restored.audio_latents.dtype == original.audio_latents.dtype
torch.testing.assert_close(
restored.audio_latents, original.audio_latents)
def test_payload_is_json_serializable(self):
envelope = _make_typed_state().to_continuation_state()
# json.dumps must not raise — required for Dynamo RPC transport
# and HTTP client round-trip.
reserialized = json.loads(json.dumps(envelope.payload))
restored = LTX2ContinuationState.from_continuation_state(
ContinuationState(
kind=envelope.kind,
payload=reserialized,
))
assert restored.segment_index == 3
def test_bf16_audio_latents_preserved(self):
"""safetensors serialization must preserve bf16 dtype (numpy
has no bf16, so a raw-bytes path would silently promote)."""
state = LTX2ContinuationState(
segment_index=0,
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.bfloat16),
)
envelope = state.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.audio_latents is not None
assert restored.audio_latents.dtype == torch.bfloat16
torch.testing.assert_close(
restored.audio_latents, state.audio_latents)
class TestBlobIndirection:
"""Large tensors live in the :class:`BlobStore` instead of the payload."""
def test_threshold_triggers_blob_path(self):
blob_store = InMemoryBlobStore()
state = _make_typed_state()
envelope = state.to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=0,
)
assert "blob_id" in envelope.payload["video"]
assert "blob_id" in envelope.payload["audio"]
assert "frames_b64" not in envelope.payload["video"]
assert "safetensors_b64" not in envelope.payload["audio"]
assert len(blob_store) == 2
def test_blob_roundtrip_reconstructs_tensors(self):
blob_store = InMemoryBlobStore()
original = _make_typed_state()
envelope = original.to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=0,
)
restored = LTX2ContinuationState.from_continuation_state(
envelope, blob_store=blob_store)
assert restored.video_frames is not None
assert len(restored.video_frames) == len(original.video_frames)
torch.testing.assert_close(
restored.audio_latents, original.audio_latents)
def test_blob_id_held_when_store_unavailable(self):
"""Deserializing without a blob store preserves the blob id so
the caller can fetch it later."""
blob_store = InMemoryBlobStore()
envelope = _make_typed_state().to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=0,
)
blob_id_video = envelope.payload["video"]["blob_id"]
blob_id_audio = envelope.payload["audio"]["blob_id"]
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.video_frames is None
assert restored.video_frames_blob_id == blob_id_video
assert restored.audio_latents is None
assert restored.audio_latents_blob_id == blob_id_audio
def test_large_threshold_keeps_payload_inline(self):
blob_store = InMemoryBlobStore()
envelope = _make_typed_state().to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=10 * 1024 * 1024, # 10 MiB
)
assert "frames_b64" in envelope.payload["video"]
assert "safetensors_b64" in envelope.payload["audio"]
assert len(blob_store) == 0
class TestValidation:
"""Invalid payloads error cleanly."""
def test_wrong_kind_rejected(self):
envelope = ContinuationState(kind="longcat.v1", payload={})
with pytest.raises(ValueError, match="Expected ContinuationState.kind"):
LTX2ContinuationState.from_continuation_state(envelope)
def test_unsupported_schema_version_rejected(self):
envelope = ContinuationState(
kind=LTX2_CONTINUATION_KIND,
payload={"schema_version": 999},
)
with pytest.raises(ValueError,
match="Unsupported LTX-2 continuation schema"):
LTX2ContinuationState.from_continuation_state(envelope)
def test_non_png_frame_rejected(self):
state = LTX2ContinuationState(
video_frames=[np.ones((64, 64, 3), dtype=np.float32)],
)
with pytest.raises(ValueError, match="uint8 HxWx3"):
state.to_continuation_state()
class TestCompatLayerWireUp:
"""The public compat layer accepts request.state without reverting
to NotImplementedError and attaches it to the SamplingParam path."""
def test_request_with_state_passes_through(self, tmp_path):
# PR 7 removes the NotImplementedError for request.state; build a
# minimal GenerationRequest carrying an LTX-2 state and make sure
# the public boundary accepts it.
from fastvideo.api.compat import (
normalize_generation_request,
_validate_continuation_state,
)
envelope = _make_typed_state().to_continuation_state()
request = GenerationRequest(
prompt="test",
state=envelope,
)
normalized = normalize_generation_request(request)
_validate_continuation_state(normalized.state)
def test_unknown_kind_rejected_at_boundary(self):
from fastvideo.api.compat import _validate_continuation_state
with pytest.raises(ValueError, match="Unknown ContinuationState kind"):
_validate_continuation_state(
ContinuationState(kind="mystery.v1", payload={}))
def test_empty_kind_rejected_at_boundary(self):
from fastvideo.api.compat import _validate_continuation_state
with pytest.raises(ValueError, match="non-empty string"):
_validate_continuation_state(
ContinuationState(kind="", payload={}))
def test_output_return_state_flag(self):
request = GenerationRequest(
prompt="x",
output=OutputConfig(return_state=True),
)
# The typed public surface exposes the flag directly.
assert request.output.return_state is True
@@ -0,0 +1,298 @@
# SPDX-License-Identifier: Apache-2.0
"""gpu_pool-style flat-kwarg integration tests.
Mirrors the ``load_kwargs`` dict that the FastVideo-internal
``ui/ltx2-streaming/server/gpu_pool.py`` passes to
``VideoGenerator.from_pretrained(**load_kwargs)`` and asserts that the
public typed ``GeneratorConfig`` surface (introduced across PRs 0-6)
can represent it end-to-end, with no fields silently falling through
to ``pipeline.experimental``.
This is the parity guard PR 7.6 depends on: the public gpu_pool
upstream must be able to construct a typed ``GeneratorConfig`` without
knowing any legacy LTX-2 kwarg name, and downstream Dynamo
(``FastVideoArgGroup``) must be able to do the same.
"""
from __future__ import annotations
from copy import deepcopy
import pytest
from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
legacy_from_pretrained_to_config,
)
# Mirrors FastVideo-internal/ui/ltx2-streaming/server/gpu_pool.py
# :lines 233-260 (load_kwargs constructed for VideoGenerator.from_pretrained).
#
# One item from gpu_pool.py's load_kwargs is deliberately excluded:
# - ``pipeline_config=<PipelineConfig instance>`` — an opaque Python
# object; internal mutates it in place (``dit_config.quant_config =
# FP4Config()``). The typed path for quantization is tracked in
# "Known Technical Debt" in PR plan.md; ``pipeline_config`` as an
# instance legitimately belongs in ``pipeline.experimental``.
#
# ``enable_torch_compile_text_encoder`` IS included below: its typed
# home is ``CompileConfig.text_encoder_enabled`` (added post-review).
# The legacy ``FastVideoArgs`` path does not yet consume it; the
# realtime runtime (PR 7.6) reads it off the kwargs dict before
# FastVideoArgs filtering.
GPU_POOL_LOAD_KWARGS = {
"config_model_path": "/models/ltx2-distilled/config",
"num_gpus": 1,
"dit_layerwise_offload": False,
"use_fsdp_inference": False,
"dit_cpu_offload": False,
"vae_cpu_offload": False,
"text_encoder_cpu_offload": False,
"pin_cpu_memory": True,
"ltx2_vae_tiling": False,
"ltx2_refine_enabled": True,
"ltx2_refine_upsampler_path": "/models/ltx2-distilled/spatial_upsampler",
"ltx2_refine_lora_path": "",
"ltx2_refine_num_inference_steps": 2,
"ltx2_refine_guidance_scale": 1.0,
"ltx2_refine_add_noise": True,
"enable_torch_compile": True,
"enable_torch_compile_text_encoder": True,
"torch_compile_kwargs": {
"backend": "inductor",
"fullgraph": True,
"mode": "max-autotune-no-cudagraphs",
"dynamic": False,
},
}
class TestGpuPoolForwardTranslation:
"""gpu_pool flat kwargs -> typed GeneratorConfig."""
@pytest.fixture(scope="class")
def config(self):
return legacy_from_pretrained_to_config(
"FastVideo/LTX2-Distilled-Diffusers",
GPU_POOL_LOAD_KWARGS,
)
def test_model_path_set(self, config) -> None:
assert config.model_path == "FastVideo/LTX2-Distilled-Diffusers"
def test_engine_basics(self, config) -> None:
assert config.engine.num_gpus == 1
assert config.engine.use_fsdp_inference is False
def test_offload_config(self, config) -> None:
assert config.engine.offload.dit is False
assert config.engine.offload.dit_layerwise is False
assert config.engine.offload.vae is False
assert config.engine.offload.text_encoder is False
assert config.engine.offload.pin_cpu_memory is True
def test_compile_config_typed_fields_extracted(self, config) -> None:
compile_config = config.engine.compile
assert compile_config.enabled is True
assert compile_config.text_encoder_enabled is True
assert compile_config.backend == "inductor"
assert compile_config.fullgraph is True
assert compile_config.mode == "max-autotune-no-cudagraphs"
assert compile_config.dynamic is False
assert compile_config.extras == {}
def test_vae_tiling_routed_to_pipeline(self, config) -> None:
assert config.pipeline.vae_tiling is False
def test_config_model_path_routed_to_components(self, config) -> None:
assert config.pipeline.components.config_root == "/models/ltx2-distilled/config"
def test_refine_upsampler_routed_to_components(self, config) -> None:
assert config.pipeline.components.upsampler_weights == (
"/models/ltx2-distilled/spatial_upsampler")
def test_empty_refine_lora_becomes_none(self, config) -> None:
# gpu_pool passes "" to keep refine LoRA disabled; typed schema
# treats that as "no LoRA" rather than an empty-string path.
assert config.pipeline.components.lora_path is None
def test_refine_preset_overrides(self, config) -> None:
refine = config.pipeline.preset_overrides.get("refine", {})
assert refine == {
"enabled": True,
"num_inference_steps": 2,
"guidance_scale": 1.0,
"add_noise": True,
}
def test_no_experimental_leakage(self, config) -> None:
"""Every gpu_pool kwarg should have a typed home — nothing should
silently fall through to ``pipeline.experimental``."""
assert config.pipeline.experimental == {}
class TestGpuPoolReverseTranslation:
"""typed GeneratorConfig -> FastVideoArgs kwargs reproduces the
original gpu_pool flat-kwarg shape.
This is what lets PR 7.6 wire the public ``gpu_pool`` through
``generator_config_to_fastvideo_args`` without the runtime noticing.
"""
@pytest.fixture
def args_kwargs(self, monkeypatch):
from fastvideo import fastvideo_args as fva
captured: dict[str, object] = {}
def _capture(**kw):
captured.update(kw)
return _Captured(**kw)
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
config = legacy_from_pretrained_to_config(
"FastVideo/LTX2-Distilled-Diffusers",
GPU_POOL_LOAD_KWARGS,
)
generator_config_to_fastvideo_args(config)
return captured
def test_ltx2_refine_flags_reemitted(self, args_kwargs) -> None:
assert args_kwargs["ltx2_refine_enabled"] is True
assert args_kwargs["ltx2_refine_add_noise"] is True
assert args_kwargs["ltx2_refine_num_inference_steps"] == 2
assert args_kwargs["ltx2_refine_guidance_scale"] == 1.0
def test_refine_upsampler_path_reemitted(self, args_kwargs) -> None:
assert args_kwargs["ltx2_refine_upsampler_path"] == (
"/models/ltx2-distilled/spatial_upsampler")
def test_config_model_path_reemitted(self, args_kwargs) -> None:
assert args_kwargs["config_model_path"] == "/models/ltx2-distilled/config"
def test_torch_compile_kwargs_reassembled(self, args_kwargs) -> None:
assert args_kwargs["torch_compile_kwargs"] == {
"backend": "inductor",
"fullgraph": True,
"mode": "max-autotune-no-cudagraphs",
"dynamic": False,
}
def test_vae_tiling_reemitted_with_legacy_name(self, args_kwargs) -> None:
assert args_kwargs["ltx2_vae_tiling"] is False
def test_text_encoder_compile_reemitted(self, args_kwargs) -> None:
# Present in the captured kwargs dict even though
# ``FastVideoArgs.from_kwargs`` will filter it out — realtime
# runtime upstream (PR 7.6) reads it off this dict.
assert args_kwargs["enable_torch_compile_text_encoder"] is True
def test_no_stray_refine_dict(self, args_kwargs) -> None:
"""preset_overrides.refine must flatten to ltx2_refine_* kwargs
rather than landing as a nested ``refine`` kwarg that
FastVideoArgs doesn't understand."""
assert "refine" not in args_kwargs
class TestRefineFlattenCoversAllTypedFields:
"""Every field on LTX2Refine{Preset,Stage}Override must survive the
round-trip through preset_overrides.refine back to ltx2_refine_*
kwargs. Guards against the hardcoded-key-tuple regression where
image_crf / video_position_offset_sec silently dropped."""
def test_all_fields_reemitted(self, monkeypatch) -> None:
from fastvideo import fastvideo_args as fva
from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
)
from fastvideo.api.schema import GeneratorConfig, PipelineSelection
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_preset_override_fields,
refine_stage_override_fields,
)
captured: dict[str, object] = {}
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
def _capture(**kw):
captured.update(kw)
return _Captured(**kw)
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
refine_payload = {
# Preset-override fields.
"enabled": True,
"add_noise": False,
# Stage-override fields.
"num_inference_steps": 3,
"guidance_scale": 1.5,
"image_crf": 18,
"video_position_offset_sec": 2.5,
}
all_fields = (refine_preset_override_fields()
| refine_stage_override_fields())
assert set(refine_payload) == all_fields, (
"payload must cover every typed field to exercise the flatten loop")
config = GeneratorConfig(
model_path="/models/ltx2",
pipeline=PipelineSelection(preset_overrides={"refine": refine_payload}),
)
generator_config_to_fastvideo_args(config)
for key, value in refine_payload.items():
assert captured[f"ltx2_refine_{key}"] == value
class TestCompileExtrasPreserved:
"""Additional torch.compile kwargs beyond the four typed fields
round-trip through ``CompileConfig.extras``."""
def test_extras_preserved(self, monkeypatch) -> None:
from fastvideo import fastvideo_args as fva
captured: dict[str, object] = {}
def _capture(**kw):
captured.update(kw)
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
return _Captured(**kw)
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
kwargs = deepcopy(GPU_POOL_LOAD_KWARGS)
kwargs["torch_compile_kwargs"] = {
"backend": "inductor",
"options": {"triton.cudagraphs": False},
"disable": False,
}
config = legacy_from_pretrained_to_config(
"FastVideo/LTX2-Distilled-Diffusers", kwargs)
assert config.engine.compile.backend == "inductor"
assert config.engine.compile.extras == {
"options": {"triton.cudagraphs": False},
"disable": False,
}
generator_config_to_fastvideo_args(config)
assert captured["torch_compile_kwargs"] == {
"backend": "inductor",
"options": {"triton.cudagraphs": False},
"disable": False,
}
@@ -0,0 +1,121 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for typed LTX-2 stage override dataclasses."""
from __future__ import annotations
import pytest
from fastvideo.api.errors import ConfigValidationError
from fastvideo.api.presets import get_preset, validate_stage_overrides
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
LTX2RefinePresetOverride,
LTX2RefineStageOverride,
refine_override_to_dict,
refine_preset_override_fields,
refine_stage_override_fields,
)
class TestRefineStageOverrideDataclass:
def test_all_fields_default_to_none(self) -> None:
override = LTX2RefineStageOverride()
assert override.num_inference_steps is None
assert override.guidance_scale is None
assert override.image_crf is None
assert override.video_position_offset_sec is None
def test_explicit_construction(self) -> None:
override = LTX2RefineStageOverride(
num_inference_steps=2,
guidance_scale=1.0,
image_crf=18,
video_position_offset_sec=2.5,
)
assert override.num_inference_steps == 2
assert override.guidance_scale == 1.0
assert override.image_crf == 18
assert override.video_position_offset_sec == 2.5
def test_to_dict_drops_none(self) -> None:
override = LTX2RefineStageOverride(num_inference_steps=3)
assert refine_override_to_dict(override) == {
"num_inference_steps": 3,
}
def test_to_dict_with_all_fields(self) -> None:
override = LTX2RefineStageOverride(
num_inference_steps=2,
guidance_scale=1.0,
image_crf=18,
video_position_offset_sec=0.0,
)
assert refine_override_to_dict(override) == {
"num_inference_steps": 2,
"guidance_scale": 1.0,
"image_crf": 18,
"video_position_offset_sec": 0.0,
}
def test_fields_accessor_matches_dataclass(self) -> None:
assert refine_stage_override_fields() == frozenset({
"num_inference_steps",
"guidance_scale",
"image_crf",
"video_position_offset_sec",
})
class TestRefinePresetOverrideDataclass:
def test_all_fields_default_to_none(self) -> None:
override = LTX2RefinePresetOverride()
assert override.enabled is None
assert override.add_noise is None
def test_to_dict_drops_none(self) -> None:
override = LTX2RefinePresetOverride(enabled=True)
assert refine_override_to_dict(override) == {
"enabled": True,
}
def test_to_dict_with_all_fields(self) -> None:
override = LTX2RefinePresetOverride(enabled=True, add_noise=False)
assert refine_override_to_dict(override) == {
"enabled": True,
"add_noise": False,
}
def test_fields_accessor_matches_dataclass(self) -> None:
assert refine_preset_override_fields() == frozenset({
"enabled",
"add_noise",
})
class TestStageOverridesMirrorPresetSchema:
"""The ltx2_two_stage preset's refine stage schema must list
exactly the :class:`LTX2RefineStageOverride` field names."""
def test_allowed_overrides_mirror_dataclass(self) -> None:
import fastvideo.registry # noqa: F401
preset = get_preset("ltx2_two_stage", "ltx2")
refine_schema = next(
s for s in preset.stage_schemas if s.name == "refine")
assert refine_schema.allowed_overrides == refine_stage_override_fields()
def test_roundtrip_through_validate_stage_overrides(self) -> None:
import fastvideo.registry # noqa: F401
preset = get_preset("ltx2_two_stage", "ltx2")
override = LTX2RefineStageOverride(
num_inference_steps=3,
guidance_scale=1.0,
)
validate_stage_overrides(
preset, {"refine": refine_override_to_dict(override)})
def test_unknown_field_rejected(self) -> None:
import fastvideo.registry # noqa: F401
preset = get_preset("ltx2_two_stage", "ltx2")
with pytest.raises(ConfigValidationError):
validate_stage_overrides(
preset, {"refine": {"unknown_key": 1}})
+10 -1
View File
@@ -111,7 +111,15 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
"vae": True,
"pin_cpu_memory": True,
},
"compile": {"enabled": False, "kwargs": {}},
"compile": {
"enabled": False,
"text_encoder_enabled": None,
"backend": None,
"fullgraph": None,
"mode": None,
"dynamic": None,
"extras": {},
},
"enable_stage_verification": True,
"use_fsdp_inference": False,
"disable_autocast": False,
@@ -133,6 +141,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
"override_pipeline_cls_name": None,
"override_transformer_cls_name": None,
},
"vae_tiling": None,
"preset_overrides": {},
"experimental": {},
},
+75 -1
View File
@@ -334,7 +334,7 @@ class TestLtx2Presets:
import fastvideo.registry # noqa: F401
presets = get_presets_for_family("ltx2")
names = {p.name for p in presets}
assert names == {"ltx2_base", "ltx2_distilled"}
assert names == {"ltx2_base", "ltx2_distilled", "ltx2_two_stage"}
def test_ltx2_base_lookup(self) -> None:
import fastvideo.registry # noqa: F401
@@ -349,6 +349,43 @@ class TestLtx2Presets:
assert p.defaults["num_inference_steps"] == 8
assert p.defaults["guidance_scale"] == 1.0
def test_ltx2_two_stage_is_two_stage(self) -> None:
import fastvideo.registry # noqa: F401
p = get_preset("ltx2_two_stage", "ltx2")
assert len(p.stage_schemas) == 2
assert p.stage_schemas[0].name == "denoise"
assert p.stage_schemas[1].name == "refine"
assert p.stage_schemas[1].kind == "refinement"
def test_ltx2_two_stage_stage_defaults(self) -> None:
import fastvideo.registry # noqa: F401
p = get_preset("ltx2_two_stage", "ltx2")
refine = p.stage_defaults["refine"]
# stage-2 refine only supports 2 or 3 denoising steps; preset
# defaults to 2 (matches gpu_pool.py load_kwargs).
assert refine["num_inference_steps"] == 2
assert refine["guidance_scale"] == 1.0
def test_ltx2_two_stage_refine_overrides_valid(self) -> None:
import fastvideo.registry # noqa: F401
p = get_preset("ltx2_two_stage", "ltx2")
validate_stage_overrides(
p, {"refine": {"num_inference_steps": 3}})
validate_stage_overrides(
p, {"refine": {"guidance_scale": 1.0}})
validate_stage_overrides(
p, {"refine": {"image_crf": 18}})
validate_stage_overrides(
p, {"refine": {"video_position_offset_sec": 2.5}})
def test_ltx2_two_stage_rejects_unknown_refine_override(self) -> None:
import fastvideo.registry # noqa: F401
from fastvideo.api.errors import ConfigValidationError
p = get_preset("ltx2_two_stage", "ltx2")
with pytest.raises(ConfigValidationError):
validate_stage_overrides(
p, {"refine": {"bogus_field": 1}})
# -------------------------------------------------------------------
# Hunyuan preset integration
@@ -548,3 +585,40 @@ class TestPresetCountIntegrity:
import fastvideo.registry # noqa: F401
names = get_all_preset_names()
assert len(names) >= 37
class TestPresetDefaultTypes:
"""Preset ``defaults`` values must match the types on
:class:`SamplingParam`. Assigning ``None`` to a typed-``str`` field
(e.g. ``negative_prompt``) breaks downstream stages that assert the
runtime type — see the CFG branch in
``pipelines/stages/text_encoding.py:81``."""
def test_ltx2_cfg_defaults_are_off(self) -> None:
"""SamplingParam's LTX-2 CFG class defaults must be 1.0 (CFG
off). ``ForwardBatch.__post_init__`` force-enables
``do_classifier_free_guidance`` when either
``ltx2_cfg_scale_video`` or ``ltx2_cfg_scale_audio`` is != 1.0,
so any non-1.0 default silently forces CFG on for every model
family that doesn't explicitly override these fields. Guard
against the regression that surfaced as the TurboDiffusion I2V
SSIM crash (``text_encoding.py:81`` assertion on
``negative_prompt``)."""
from fastvideo.api.sampling_param import SamplingParam
sp = SamplingParam()
assert sp.ltx2_cfg_scale_video == 1.0
assert sp.ltx2_cfg_scale_audio == 1.0
def test_no_preset_sets_negative_prompt_to_none(self) -> None:
import fastvideo.registry # noqa: F401
from fastvideo.api.presets import _PRESET_REGISTRY
offenders = [
f"{preset.model_family}/{preset.name}"
for preset in _PRESET_REGISTRY.values()
if preset.defaults.get("negative_prompt", "") is None
]
assert not offenders, (
"These presets set negative_prompt=None, which violates "
"SamplingParam.negative_prompt's typed str contract and "
"crashes the CFG path in text_encoding. Use \"\" instead:\n"
+ "\n".join(f" - {p}" for p in offenders))
@@ -46,24 +46,42 @@ def _flatten_status_section(section: dict, valid_statuses: set[str]) -> set[str]
return names
def _get_extra_dataclass_fields(package_name: str, base_cls: type) -> set[str]:
package = importlib.import_module(package_name)
def _get_extra_dataclass_fields(
package_names: str | tuple[str, ...],
base_cls: type,
) -> set[str]:
"""Collect dataclass fields declared on ``base_cls`` subclasses found
under any of the given package roots.
Accepts either a single package name (string) or a tuple of package
roots — the latter supports the PR 6 colocation where each model
family's ``PipelineConfig`` subclass moves from
``fastvideo.configs.pipelines.<family>`` to
``fastvideo.pipelines.basic.<family>.pipeline_configs``.
"""
if isinstance(package_names, str):
package_names = (package_names, )
base_fields = {f.name for f in dataclasses.fields(base_cls)}
extras: set[str] = set()
if not hasattr(package, "__path__"):
return extras
for _, modname, _ in pkgutil.iter_modules(package.__path__):
if modname == "__pycache__":
for package_name in package_names:
package = importlib.import_module(package_name)
if not hasattr(package, "__path__"):
continue
module = importlib.import_module(f"{package_name}.{modname}")
for obj in vars(module).values():
if (
isinstance(obj, type)
and dataclasses.is_dataclass(obj)
and issubclass(obj, base_cls)
and obj is not base_cls
):
extras.update(f.name for f in dataclasses.fields(obj) if f.name not in base_fields)
for _, modname, is_pkg in pkgutil.walk_packages(
package.__path__, prefix=f"{package_name}."):
if modname.endswith(".__pycache__"):
continue
module = importlib.import_module(modname)
for obj in vars(module).values():
if (isinstance(obj, type)
and dataclasses.is_dataclass(obj)
and issubclass(obj, base_cls)
and obj is not base_cls):
extras.update(
f.name for f in dataclasses.fields(obj)
if f.name not in base_fields)
return extras
@@ -177,7 +195,10 @@ def test_pipeline_config_base_fields_are_classified() -> None:
def test_pipeline_config_extension_fields_are_classified() -> None:
inventory = _load_inventory()
expected = _get_extra_dataclass_fields("fastvideo.configs.pipelines", PipelineConfig)
expected = _get_extra_dataclass_fields(
("fastvideo.configs.pipelines", "fastvideo.pipelines.basic"),
PipelineConfig,
)
actual = _flatten_status_section(
inventory["surfaces"]["pipeline_config_extensions"],
set(inventory["status_definitions"]),
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
"""Protocol schema tests for the streaming server.
Covers:
* accepted client messages parse into the correct discriminated model
* unknown ``type`` values raise validation errors
* server-side messages serialize to the expected wire shape
* continuation_state field on session_init_v2 carries through
"""
from __future__ import annotations
import pytest
from pydantic import ValidationError
from fastvideo.entrypoints.streaming.protocol import (
ContinuationStateSnapshot,
ErrorMessage,
GpuAssigned,
Ltx2SegmentComplete,
Ltx2SegmentStart,
Ltx2StreamStart,
MediaInit,
MediaSegmentComplete,
QueueStatus,
SegmentPromptSource,
SessionInitV2,
SnapshotState,
StepComplete,
parse_client_message,
)
class TestClientMessageParsing:
def test_session_init_v2_minimal(self):
parsed = parse_client_message({"type": "session_init_v2"})
assert isinstance(parsed, SessionInitV2)
assert parsed.curated_prompts == []
assert parsed.stream_mode == "av_fmp4"
def test_session_init_v2_full(self):
raw = {
"type": "session_init_v2",
"client_id": "client-1",
"preset": "ltx2_two_stage",
"preset_label": "2x refine",
"curated_prompts": ["a fox", "a deer"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
"single_clip_mode": True,
"stream_mode": "av_fmp4",
"continuation_state": {
"kind": "ltx2.v1",
"payload": {"schema_version": 1, "segment_index": 2},
},
}
parsed = parse_client_message(raw)
assert isinstance(parsed, SessionInitV2)
assert parsed.preset == "ltx2_two_stage"
assert parsed.curated_prompts == ["a fox", "a deer"]
assert parsed.continuation_state["kind"] == "ltx2.v1"
def test_segment_prompt_source(self):
parsed = parse_client_message({
"type": "segment_prompt_source",
"prompt": "hello world",
"source": "curated",
"seed": 7,
})
assert isinstance(parsed, SegmentPromptSource)
assert parsed.source == "curated"
assert parsed.seed == 7
def test_snapshot_state(self):
parsed = parse_client_message({"type": "snapshot_state"})
assert isinstance(parsed, SnapshotState)
def test_unknown_type_rejected(self):
with pytest.raises(ValidationError):
parse_client_message({"type": "not_a_real_message"})
def test_missing_type_rejected(self):
with pytest.raises(ValidationError):
parse_client_message({"prompt": "x"})
def test_segment_prompt_source_requires_prompt(self):
with pytest.raises(ValidationError):
parse_client_message({"type": "segment_prompt_source"})
class TestServerMessageSerialization:
def test_queue_status(self):
msg = QueueStatus(position=3, queue_depth=5)
assert msg.model_dump() == {
"type": "queue_status",
"position": 3,
"queue_depth": 5,
}
def test_gpu_assigned(self):
msg = GpuAssigned(gpu_id=1, session_timeout=300)
assert msg.model_dump()["type"] == "gpu_assigned"
def test_ltx2_stream_start(self):
msg = Ltx2StreamStart(
preset="ltx2_two_stage",
width=1024, height=1536, fps=24, num_frames=121,
)
dumped = msg.model_dump()
assert dumped["type"] == "ltx2_stream_start"
assert dumped["width"] == 1024
def test_ltx2_segment_start(self):
msg = Ltx2SegmentStart(
segment_idx=0,
prompt="a fox",
total_steps=8,
)
assert msg.model_dump()["segment_idx"] == 0
def test_step_complete(self):
msg = StepComplete(segment_idx=0, step=1, total_steps=8)
assert msg.model_dump()["stage"] == "denoise"
def test_media_init_has_mode(self):
msg = MediaInit(segment_idx=0, stream_id="abc")
dumped = msg.model_dump()
assert dumped["mode"] == "av_fmp4"
assert "avc1" in dumped["mime"]
def test_media_segment_complete(self):
msg = MediaSegmentComplete(
segment_idx=0, stream_id="abc", chunks=4,
)
dumped = msg.model_dump()
assert dumped["chunks"] == 4
def test_ltx2_segment_complete(self):
msg = Ltx2SegmentComplete(segment_idx=0, generation_time_ms=1234.5)
assert msg.model_dump()["generation_time_ms"] == 1234.5
def test_error_message_code_restricted(self):
with pytest.raises(ValidationError):
ErrorMessage(code="not_a_code", message="x")
def test_continuation_state_snapshot(self):
msg = ContinuationStateSnapshot(state={
"kind": "ltx2.v1",
"payload": {"schema_version": 1},
})
assert msg.model_dump()["state"]["kind"] == "ltx2.v1"
@@ -0,0 +1,237 @@
# SPDX-License-Identifier: Apache-2.0
"""End-to-end WebSocket smoke for the streaming server skeleton.
Uses a mock generator so these tests run CPU-only (no GPU, no model
weights). Skips the fMP4 assertions when ``ffmpeg`` is missing.
"""
from __future__ import annotations
import shutil
from dataclasses import dataclass
from typing import Any
import numpy as np
import pytest
pytest.importorskip("starlette")
from starlette.testclient import TestClient # noqa: E402
from fastvideo.api.schema import ( # noqa: E402
ContinuationState,
GeneratorConfig,
SamplingConfig,
ServeConfig,
StreamingConfig,
GenerationRequest,
)
from fastvideo.entrypoints.streaming.server import build_app # noqa: E402
_FFMPEG_AVAILABLE = shutil.which("ffmpeg") is not None
@dataclass
class _MockGenerator:
width: int = 64
height: int = 64
fps: int = 12
num_frames: int = 12
return_state: bool = True
def generate(self, request: GenerationRequest) -> dict[str, Any]:
frames = [
np.full((self.height, self.width, 3), i * 5, dtype=np.uint8)
for i in range(self.num_frames)
]
state = (ContinuationState(
kind="ltx2.v1",
payload={
"schema_version": 1,
"segment_index": 0,
"source_prompt": request.prompt,
},
) if self.return_state else None)
return {
"frames": frames,
"audio_sample_rate": 24000,
"state": state,
}
def _build_serve_config() -> ServeConfig:
return ServeConfig(
generator=GeneratorConfig(model_path="/models/fake"),
default_request=GenerationRequest(
sampling=SamplingConfig(
num_frames=12,
height=64,
width=64,
fps=12,
num_inference_steps=1,
),
),
streaming=StreamingConfig(
session_timeout_seconds=60,
generation_segment_cap=2,
),
)
def _build_client() -> tuple[TestClient, _MockGenerator]:
generator = _MockGenerator()
app = build_app(_build_serve_config(), generator)
return TestClient(app), generator
class TestHealth:
def test_health_endpoint_reports_stream_mode(self):
client, _ = _build_client()
response = client.get("/health")
assert response.status_code == 200
body = response.json()
assert body["status"] == "ok"
assert body["stream_mode"] == "av_fmp4"
assert body["sessions"] == 0
class TestSessionHandshake:
def test_rejects_non_session_init_opening_frame(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "segment_prompt_source", "prompt": "x"})
err = ws.receive_json()
assert err["type"] == "error"
assert err["code"] == "invalid_message"
def test_rejects_unknown_message_on_init(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "not_a_message"})
err = ws.receive_json()
assert err["type"] == "error"
def test_emits_queue_and_gpu_assigned_on_valid_init(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({
"type": "session_init_v2",
"preset": "ltx2_two_stage",
"curated_prompts": ["a fox"],
})
assert ws.receive_json()["type"] == "queue_status"
assert ws.receive_json()["type"] == "gpu_assigned"
assert ws.receive_json()["type"] == "ltx2_stream_start"
def test_init_hydrates_continuation_state(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({
"type": "session_init_v2",
"preset": "ltx2_two_stage",
"continuation_state": {
"kind": "ltx2.v1",
"payload": {"schema_version": 1, "segment_index": 3},
},
})
# Drain handshake frames
ws.receive_json() # queue_status
ws.receive_json() # gpu_assigned
ws.receive_json() # ltx2_stream_start
# Ask the server for the state back; it should echo what we sent.
ws.send_json({"type": "snapshot_state"})
snap = ws.receive_json()
assert snap["type"] == "continuation_state_snapshot"
assert snap["state"]["kind"] == "ltx2.v1"
assert snap["state"]["payload"]["segment_index"] == 3
@pytest.mark.skipif(not _FFMPEG_AVAILABLE, reason="ffmpeg not installed")
class TestSegmentFlow:
def test_segment_generates_media_init_plus_complete(self):
client, generator = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "session_init_v2",
"preset": "ltx2_two_stage"})
for _ in range(3):
ws.receive_json() # queue_status + gpu_assigned + stream_start
ws.send_json({
"type": "segment_prompt_source",
"prompt": "a test segment",
"num_inference_steps": 1,
})
start = ws.receive_json()
assert start["type"] == "ltx2_segment_start"
assert start["segment_idx"] == 0
step = ws.receive_json()
assert step["type"] == "step_complete"
media_init = ws.receive_json()
assert media_init["type"] == "media_init"
# Then one or more binary frames until media_segment_complete.
saw_binary = False
while True:
msg = ws.receive()
if "bytes" in msg and msg["bytes"]:
saw_binary = True
continue
parsed = _as_json(msg)
if parsed is None:
continue
if parsed["type"] == "media_segment_complete":
break
assert saw_binary
final = ws.receive_json()
assert final["type"] == "ltx2_segment_complete"
assert final["segment_idx"] == 0
class TestContinuationStatePersistence:
def test_snapshot_after_segment_carries_generator_state(self):
if not _FFMPEG_AVAILABLE:
pytest.skip("ffmpeg not installed")
client, generator = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "session_init_v2",
"preset": "ltx2_two_stage"})
for _ in range(3):
ws.receive_json()
ws.send_json({
"type": "segment_prompt_source",
"prompt": "a cat",
"num_inference_steps": 1,
})
_drain_until(ws, "ltx2_segment_complete")
ws.send_json({"type": "snapshot_state"})
snap = ws.receive_json()
assert snap["type"] == "continuation_state_snapshot"
assert snap["state"]["kind"] == "ltx2.v1"
assert snap["state"]["payload"]["source_prompt"] == "a cat"
# ----------------------------------------------------------------------
# Helpers
# ----------------------------------------------------------------------
def _drain_until(ws, target_type: str) -> dict[str, Any]:
while True:
msg = ws.receive()
if "text" in msg and msg["text"]:
import json
parsed = json.loads(msg["text"])
if parsed.get("type") == target_type:
return parsed
# skip binary / other
def _as_json(msg: dict[str, Any]) -> dict[str, Any] | None:
if "text" not in msg or not msg["text"]:
return None
import json
return json.loads(msg["text"])
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
"""Session lifecycle tests."""
from __future__ import annotations
import time
import pytest
from fastvideo.entrypoints.streaming.session import (
InvalidSessionTransition,
Session,
SessionManager,
SessionRejected,
SessionState,
)
class TestSessionStateMachine:
def test_starts_initializing(self):
s = Session()
assert s.state is SessionState.INITIALIZING
def test_legal_sequence(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.transition(SessionState.COMPLETE)
assert s.state is SessionState.COMPLETE
def test_active_self_loop_allowed(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.transition(SessionState.ACTIVE) # re-asserting is fine
assert s.state is SessionState.ACTIVE
def test_illegal_backwards_transition(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
with pytest.raises(InvalidSessionTransition):
s.transition(SessionState.INITIALIZING)
def test_cannot_leave_terminal_state(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.transition(SessionState.COMPLETE)
with pytest.raises(InvalidSessionTransition):
s.transition(SessionState.ACTIVE)
def test_error_terminal(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.ERROR)
with pytest.raises(InvalidSessionTransition):
s.transition(SessionState.ACTIVE)
def test_transition_updates_activity(self):
s = Session()
prior = s.last_activity
time.sleep(0.001)
s.transition(SessionState.QUEUED)
assert s.last_activity > prior
def test_segment_cap(self):
s = Session()
s.segment_idx = 5
assert s.segment_cap_reached(5) is True
assert s.segment_cap_reached(6) is False
class TestSessionManager:
def test_create_assigns_unique_ids(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=2)
a = mgr.create()
b = mgr.create()
assert a.id != b.id
assert len(mgr) == 2
def test_max_sessions_enforced(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
mgr.create()
with pytest.raises(SessionRejected):
mgr.create()
def test_close_releases_slot(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
s = mgr.create()
mgr.close(s.id)
assert len(mgr) == 0
# Now can create again.
mgr.create()
def test_reap_timed_out_flags_idle_sessions(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
s = mgr.create()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.last_activity = time.monotonic() - 10 # 10s ago, past the 1s budget
dead = mgr.reap_timed_out()
assert s.id in dead
def test_reap_skips_terminal_states(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
s = mgr.create()
s.transition(SessionState.QUEUED)
s.transition(SessionState.ERROR)
s.last_activity = time.monotonic() - 10
assert s.id not in mgr.reap_timed_out()
def test_active_sessions_filter(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=4)
a = mgr.create()
a.transition(SessionState.QUEUED)
a.transition(SessionState.GPU_BINDING)
a.transition(SessionState.ACTIVE)
b = mgr.create() # INITIALIZING
assert mgr.active_sessions() == [a]
assert b not in mgr.active_sessions()
@@ -0,0 +1,90 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the session init-image persistence helper."""
from __future__ import annotations
import base64
import io
import os
import pytest
from PIL import Image
from fastvideo.entrypoints.streaming.session_init_image import (
persist_session_init_image,
)
def _png_bytes(size: tuple[int, int] = (64, 64)) -> bytes:
buffer = io.BytesIO()
Image.new("RGB", size, color=(10, 20, 30)).save(buffer, format="PNG")
return buffer.getvalue()
class TestPersistSessionInitImage:
def test_none_payload_returns_none(self):
assert persist_session_init_image(None) is None
assert persist_session_init_image({}) is None
def test_non_object_payload_rejected(self):
with pytest.raises(ValueError):
persist_session_init_image("not-a-dict")
def test_png_payload_persists(self, tmp_path):
data = _png_bytes()
image = persist_session_init_image({
"mime": "image/png",
"name": "ref.png",
"data": base64.b64encode(data).decode("ascii"),
}, output_dir=str(tmp_path))
assert image is not None
assert os.path.exists(image.path)
assert image.mime == "image/png"
assert image.path.endswith(".png")
with open(image.path, "rb") as f:
assert f.read() == data
def test_unknown_mime_rejected(self, tmp_path):
with pytest.raises(ValueError, match="mime"):
persist_session_init_image({
"mime": "image/bmp",
"data": "ignored",
}, output_dir=str(tmp_path))
def test_bad_base64_rejected(self, tmp_path):
with pytest.raises(ValueError, match="base64"):
persist_session_init_image({
"mime": "image/png",
"data": "not!base64!",
}, output_dir=str(tmp_path))
def test_empty_data_rejected(self, tmp_path):
with pytest.raises(ValueError, match="empty"):
persist_session_init_image({
"mime": "image/png",
"data": "",
}, output_dir=str(tmp_path))
def test_display_name_sanitized(self, tmp_path):
image = persist_session_init_image({
"mime": "image/png",
"name": "../evil/../name.png",
"data": base64.b64encode(_png_bytes()).decode("ascii"),
}, output_dir=str(tmp_path))
assert image is not None
assert image.display_name == "name.png"
def test_oversize_rejected(self, tmp_path):
from fastvideo.entrypoints.streaming import session_init_image as mod
original = mod._MAX_IMAGE_BYTES
mod._MAX_IMAGE_BYTES = 100
try:
with pytest.raises(ValueError, match="limit"):
persist_session_init_image({
"mime": "image/png",
"data": base64.b64encode(_png_bytes((512, 512))).decode(
"ascii"),
}, output_dir=str(tmp_path))
finally:
mod._MAX_IMAGE_BYTES = original
@@ -0,0 +1,186 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the streaming SessionStore and BlobStore.
Covers:
* ``store`` / ``snapshot`` / ``drop`` lifecycle for the in-memory store
* ``hydrate`` with and without an explicit session id
* blob store insert / get / drop semantics
* thread-safety under concurrent writes (smoke)
* round-trip a LTX-2 continuation through snapshot + hydrate across a
session boundary (the "export and resume" flow the PR plan calls out)
"""
from __future__ import annotations
import threading
import numpy as np
import pytest
import torch
from fastvideo.api.schema import ContinuationState
from fastvideo.entrypoints.streaming.session_store import (
BlobStore,
InMemoryBlobStore,
InMemorySessionStore,
SessionStore,
)
from fastvideo.pipelines.basic.ltx2.continuation import (
LTX2_CONTINUATION_KIND,
LTX2ContinuationState,
)
class TestInMemoryBlobStore:
def test_is_blob_store(self):
assert isinstance(InMemoryBlobStore(), BlobStore)
def test_put_then_get_returns_same_bytes(self):
store = InMemoryBlobStore()
blob_id = store.put(b"hello")
assert store.get(blob_id) == b"hello"
def test_put_returns_distinct_ids(self):
store = InMemoryBlobStore()
id_a = store.put(b"a")
id_b = store.put(b"b")
assert id_a != id_b
def test_get_missing_raises_keyerror(self):
store = InMemoryBlobStore()
with pytest.raises(KeyError):
store.get("nonexistent")
def test_drop_removes_blob(self):
store = InMemoryBlobStore()
blob_id = store.put(b"payload")
store.drop(blob_id)
assert blob_id not in store
with pytest.raises(KeyError):
store.get(blob_id)
def test_drop_missing_is_noop(self):
store = InMemoryBlobStore()
store.drop("not-there") # no raise
def test_contains(self):
store = InMemoryBlobStore()
blob_id = store.put(b"x")
assert blob_id in store
assert "other" not in store
class TestInMemorySessionStore:
def test_is_session_store(self):
assert isinstance(InMemorySessionStore(), SessionStore)
def test_store_then_snapshot(self):
store = InMemorySessionStore()
state = ContinuationState(kind="ltx2.v1", payload={"x": 1})
store.store("sess-1", state)
assert store.snapshot("sess-1") is state
def test_snapshot_missing_returns_none(self):
store = InMemorySessionStore()
assert store.snapshot("missing") is None
def test_store_overwrites_prior_state(self):
store = InMemorySessionStore()
first = ContinuationState(kind="ltx2.v1", payload={"v": 1})
second = ContinuationState(kind="ltx2.v1", payload={"v": 2})
store.store("s", first)
store.store("s", second)
assert store.snapshot("s").payload["v"] == 2
def test_hydrate_assigns_new_session_id(self):
store = InMemorySessionStore()
state = ContinuationState(kind="ltx2.v1", payload={})
sid = store.hydrate(state)
assert sid
assert store.snapshot(sid) is state
def test_hydrate_with_explicit_session_id(self):
store = InMemorySessionStore()
state = ContinuationState(kind="ltx2.v1", payload={})
sid = store.hydrate(state, session_id="pinned-id")
assert sid == "pinned-id"
assert store.snapshot("pinned-id") is state
def test_drop_forgets_session(self):
store = InMemorySessionStore()
store.store("s", ContinuationState(kind="ltx2.v1", payload={}))
store.drop("s")
assert store.snapshot("s") is None
assert "s" not in store
def test_iter_yields_session_ids(self):
store = InMemorySessionStore()
store.store("a", ContinuationState(kind="ltx2.v1", payload={}))
store.store("b", ContinuationState(kind="ltx2.v1", payload={}))
assert sorted(store) == ["a", "b"]
def test_len(self):
store = InMemorySessionStore()
assert len(store) == 0
store.store("x", ContinuationState(kind="ltx2.v1", payload={}))
assert len(store) == 1
def test_concurrent_store_is_safe(self):
"""Smoke-check the lock: 200 parallel stores settle to 200 ids."""
store = InMemorySessionStore()
def write(i: int) -> None:
store.store(
f"s-{i}",
ContinuationState(kind="ltx2.v1", payload={"i": i}),
)
threads = [threading.Thread(target=write, args=(i,)) for i in range(200)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(store) == 200
class TestSnapshotHydrateRoundTrip:
"""Session boundary: snapshot + hydrate preserves the full LTX-2 state."""
def test_end_to_end_ltx2_session_migration(self):
blob_store = InMemoryBlobStore()
sessions = InMemorySessionStore()
typed = LTX2ContinuationState(
segment_index=4,
video_frames=[
np.full((32, 32, 3), i * 5, dtype=np.uint8) for i in range(3)
],
audio_latents=torch.randn(1, 4, 8, 32, dtype=torch.float32),
audio_sample_rate=24000,
audio_conditioning_num_frames=5,
video_position_offset_sec=0.25,
)
envelope = typed.to_continuation_state(blob_store=blob_store)
sessions.store("session-a", envelope)
snapshot = sessions.snapshot("session-a")
assert snapshot is not None
assert snapshot.kind == LTX2_CONTINUATION_KIND
# Simulate a migration: drop the first session, hydrate a new one
# from the snapshot, and reconstruct the typed state.
sessions.drop("session-a")
new_sid = sessions.hydrate(snapshot)
assert new_sid != "session-a"
rebuilt = sessions.snapshot(new_sid)
assert rebuilt is snapshot
restored = LTX2ContinuationState.from_continuation_state(
rebuilt, blob_store=blob_store)
assert restored.segment_index == typed.segment_index
assert restored.audio_sample_rate == typed.audio_sample_rate
torch.testing.assert_close(
restored.audio_latents, typed.audio_latents)
assert len(restored.video_frames) == len(typed.video_frames)
@@ -0,0 +1,99 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the fMP4 encoder.
These tests require ``ffmpeg`` on PATH. Skip when missing so the suite
stays CPU/CI friendly.
"""
from __future__ import annotations
import asyncio
import shutil
import numpy as np
import pytest
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
)
pytestmark = pytest.mark.skipif(
shutil.which("ffmpeg") is None,
reason="ffmpeg not installed",
)
def _frame(width: int, height: int, value: int = 128) -> np.ndarray:
return np.full((height, width, 3), value, dtype=np.uint8)
def test_encoder_emits_init_then_media_chunks():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
chunks: list[FragmentedMP4Chunk] = []
async with enc:
frames = [_frame(64, 64, v) for v in range(4, 28)]
async for chunk in enc.encode(frames):
chunks.append(chunk)
assert len(chunks) > 0
assert chunks[0].kind == "init"
assert all(c.stream_id == enc.stream_id for c in chunks)
assert all(c.segment_idx == 0 for c in chunks)
asyncio.run(run())
def test_encoder_init_chunk_is_fmp4():
"""The first chunk must contain the ``ftyp`` box (fMP4 init segment)."""
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
first_chunk = None
async with enc:
async for chunk in enc.encode([_frame(64, 64, 20)] * 24):
first_chunk = chunk
break
assert first_chunk is not None
assert first_chunk.kind == "init"
# Box header: 4 bytes length, 4 bytes type. "ftyp" should appear
# near the start of the init segment.
assert b"ftyp" in first_chunk.data[:32]
asyncio.run(run())
def test_encoder_rejects_non_ndarray_frames():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
async with enc:
with pytest.raises(TypeError):
async for _ in enc.encode(["not-a-frame"]):
pass
asyncio.run(run())
def test_encoder_rejects_wrong_shape():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
async with enc:
with pytest.raises(ValueError):
async for _ in enc.encode(
[np.zeros((64, 64, 4), dtype=np.uint8)]):
pass
asyncio.run(run())
def test_encoder_close_is_idempotent():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
await enc.__aenter__()
await enc.close()
await enc.close() # no raise
asyncio.run(run())
+1 -1
View File
@@ -480,10 +480,10 @@ def _prepare_ssim_workspace(
{checkout_command}
rm -rf fastvideo/tests/ssim/reference_videos
git_retry git submodule update --init --recursive
uv pip install -e .[test]
cd fastvideo-kernel
./build.sh
cd ..
uv pip install -e .[test]
uv pip install git+https://github.com/microsoft/MoGe.git
export HF_HOME='/root/data/.cache'
hf auth login --token "$HF_API_KEY"
@@ -23,6 +23,7 @@ DEVICE_MAPPINGS = (
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
("B200", "B200"),
)
+52 -3
View File
@@ -251,6 +251,8 @@ def upload_reference_videos(
reference_dirs_by_tier: Sequence[tuple[str, Path]],
token: str,
private: bool,
model_id: str | None = None,
force: bool = False,
) -> None:
HfApi, _ = _load_hf_sdk()
api = HfApi(token=token)
@@ -261,15 +263,44 @@ def upload_reference_videos(
exist_ok=True,
)
try:
existing_repo_files = set(
api.list_repo_files(repo_id=repo_id, repo_type=repo_type))
except Exception:
# Fresh repo or list failure — treat as empty so upload can proceed.
existing_repo_files = set()
for quality_tier, reference_dir in reference_dirs_by_tier:
if not reference_dir.exists():
raise FileNotFoundError(f"Reference directory not found: {reference_dir}")
path_in_repo = f"{REFERENCE_VIDEOS_DIRNAME}/{quality_tier}/{reference_dir.name}"
print(f"Uploading {reference_dir.name} ({quality_tier}) to {repo_id} ...")
base_in_repo = f"{REFERENCE_VIDEOS_DIRNAME}/{quality_tier}/{reference_dir.name}"
if model_id:
folder_path = reference_dir / model_id
if not folder_path.exists():
raise FileNotFoundError(
f"Model subfolder not found for upload: {folder_path}")
path_in_repo = f"{base_in_repo}/{model_id}"
else:
folder_path = reference_dir
path_in_repo = base_in_repo
conflicts = sorted(
f for f in existing_repo_files
if f.startswith(f"{path_in_repo}/") or f == path_in_repo)
if conflicts and not force:
preview = "\n".join(f" - {c}" for c in conflicts[:10])
more = f"\n ... and {len(conflicts) - 10} more" if len(conflicts) > 10 else ""
raise RuntimeError(
f"Refusing to overwrite existing HF files under {path_in_repo} "
f"({len(conflicts)} file(s) already present):\n{preview}{more}\n"
f"Re-run with --force to overwrite.")
target_desc = f"{reference_dir.name}/{model_id}" if model_id else reference_dir.name
print(f"Uploading {target_desc} ({quality_tier}) to {repo_id}/{path_in_repo} ...")
api.upload_folder(
repo_id=repo_id,
repo_type=repo_type,
folder_path=str(reference_dir),
folder_path=str(folder_path),
path_in_repo=path_in_repo,
token=token,
)
@@ -500,6 +531,22 @@ def _build_parser() -> argparse.ArgumentParser:
action="store_true",
help="Create/use a private repo instead of public.",
)
upload_parser.add_argument(
"--model-id",
default=None,
help=(
"Restrict upload to a single model subfolder "
"(reference_videos/<tier>/<device>/<model_id>). "
"Use when seeding references for a single new test."),
)
upload_parser.add_argument(
"--force",
action="store_true",
help=(
"Allow overwriting files that already exist at the target path on "
"Hugging Face. Off by default so seeding a new test cannot "
"clobber existing references."),
)
ensure_parser = subparsers.add_parser(
"ensure",
@@ -607,6 +654,8 @@ def main(argv: Sequence[str] | None = None) -> int:
reference_dirs_by_tier=reference_dirs_by_tier,
token=token,
private=args.private,
model_id=args.model_id,
force=args.force,
)
print("Upload complete.")
return 0
+110 -110
View File
@@ -10,21 +10,21 @@ Note: num_inference_steps is reduced to 4 for faster CI.
import os
import pytest
import torch
from fastvideo import VideoGenerator
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.logger import init_logger
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
from fastvideo.tests.ssim.reference_utils import (
build_generated_output_dir,
build_reference_folder_path,
get_cuda_device_name,
resolve_device_reference_folder,
select_ssim_params,
)
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
import pytest
import torch
from fastvideo import VideoGenerator
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.logger import init_logger
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
from fastvideo.tests.ssim.reference_utils import (
build_generated_output_dir,
build_reference_folder_path,
get_cuda_device_name,
resolve_device_reference_folder,
select_ssim_params,
)
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
logger = init_logger(__name__)
@@ -48,20 +48,20 @@ def _find_lingbotworld_examples_root() -> str | None:
return None
device_name = get_cuda_device_name()
device_reference_folder = resolve_device_reference_folder(
(
("A40", "A40"),
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
),
device_name=device_name,
logger=logger,
)
device_name = get_cuda_device_name()
device_reference_folder = resolve_device_reference_folder(
(
("A40", "A40"),
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
),
device_name=device_name,
logger=logger,
)
LINGBOT_PARAMS = {
LINGBOT_PARAMS = {
"model_path": "FastVideo/LingBot-World-Base-Cam-Diffusers",
"num_gpus": 2,
"height": 256,
@@ -88,29 +88,29 @@ LINGBOT_PARAMS = {
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
"皮肤,肢体,面部特征,汽车,电线"
),
}
_LINGBOT_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
LINGBOT_PARAMS["model_path"])
LINGBOT_FULL_QUALITY_PARAMS = {
"model_path": LINGBOT_PARAMS["model_path"],
"num_gpus": LINGBOT_PARAMS["num_gpus"],
"height": _LINGBOT_FULL_QUALITY_DEFAULTS.height,
"width": _LINGBOT_FULL_QUALITY_DEFAULTS.width,
"num_frames": LINGBOT_PARAMS["num_frames"], # default num_frames: 125
"num_inference_steps": _LINGBOT_FULL_QUALITY_DEFAULTS.num_inference_steps,
"guidance_scale": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale,
"guidance_scale_2": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale_2,
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
"flow_shift": LINGBOT_PARAMS["flow_shift"],
"boundary_ratio": _LINGBOT_FULL_QUALITY_DEFAULTS.boundary_ratio,
"seed": _LINGBOT_FULL_QUALITY_DEFAULTS.seed,
"fps": _LINGBOT_FULL_QUALITY_DEFAULTS.fps,
"spatial_scale": LINGBOT_PARAMS["spatial_scale"],
"example_case": LINGBOT_PARAMS["example_case"],
"image_path": LINGBOT_PARAMS["image_path"],
"negative_prompt": _LINGBOT_FULL_QUALITY_DEFAULTS.negative_prompt,
}
),
}
_LINGBOT_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
LINGBOT_PARAMS["model_path"])
LINGBOT_FULL_QUALITY_PARAMS = {
"model_path": LINGBOT_PARAMS["model_path"],
"num_gpus": LINGBOT_PARAMS["num_gpus"],
"height": _LINGBOT_FULL_QUALITY_DEFAULTS.height,
"width": _LINGBOT_FULL_QUALITY_DEFAULTS.width,
"num_frames": LINGBOT_PARAMS["num_frames"], # default num_frames: 125
"num_inference_steps": _LINGBOT_FULL_QUALITY_DEFAULTS.num_inference_steps,
"guidance_scale": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale,
"guidance_scale_2": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale_2,
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
"flow_shift": LINGBOT_PARAMS["flow_shift"],
"boundary_ratio": _LINGBOT_FULL_QUALITY_DEFAULTS.boundary_ratio,
"seed": _LINGBOT_FULL_QUALITY_DEFAULTS.seed,
"fps": _LINGBOT_FULL_QUALITY_DEFAULTS.fps,
"spatial_scale": LINGBOT_PARAMS["spatial_scale"],
"example_case": LINGBOT_PARAMS["example_case"],
"image_path": LINGBOT_PARAMS["image_path"],
"negative_prompt": _LINGBOT_FULL_QUALITY_DEFAULTS.negative_prompt,
}
TEST_PROMPTS = [
"The video presents a soaring journey through a fantasy jungle. The wind "
@@ -123,80 +123,80 @@ TEST_PROMPTS = [
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
params = select_ssim_params(LINGBOT_PARAMS, LINGBOT_FULL_QUALITY_PARAMS)
if device_reference_folder is None:
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
if torch.cuda.device_count() < params["num_gpus"]:
pytest.skip(
f"LingBot SSIM test requires {params['num_gpus']} GPUs, "
f"but only {torch.cuda.device_count()} detected."
)
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
params = select_ssim_params(LINGBOT_PARAMS, LINGBOT_FULL_QUALITY_PARAMS)
if device_reference_folder is None:
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
if torch.cuda.device_count() < params["num_gpus"]:
pytest.skip(
f"LingBot SSIM test requires {params['num_gpus']} GPUs, "
f"but only {torch.cuda.device_count()} detected."
)
examples_root = _find_lingbotworld_examples_root()
if examples_root is None:
pytest.skip(
"lingbotworld_examples not found under examples/inference/basic.")
action_path = os.path.join(examples_root, params["example_case"])
action_path = os.path.join(examples_root, params["example_case"])
if not (os.path.exists(os.path.join(action_path, "poses.npy"))
and os.path.exists(os.path.join(action_path, "intrinsics.npy"))):
pytest.skip(f"Missing camera npy files under {action_path}")
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=params["num_frames"],
height=params["height"],
width=params["width"],
spatial_scale=params["spatial_scale"],
)
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=params["num_frames"],
height=params["height"],
width=params["width"],
spatial_scale=params["spatial_scale"],
)
script_dir = os.path.dirname(os.path.abspath(__file__))
model_id = "LingBot-World-Base-Cam-Diffusers"
output_dir = build_generated_output_dir(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
output_dir = build_generated_output_dir(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
init_kwargs = {
"num_gpus": params["num_gpus"],
"flow_shift": params["flow_shift"],
"boundary_ratio": params["boundary_ratio"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
init_kwargs = {
"num_gpus": params["num_gpus"],
"flow_shift": params["flow_shift"],
"boundary_ratio": params["boundary_ratio"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": True,
"vae_cpu_offload": False,
"pin_cpu_memory": True,
}
generation_kwargs = {
"output_path": output_dir,
"image_path": params["image_path"],
"height": params["height"],
"width": params["width"],
"num_frames": aligned_num_frames,
"num_inference_steps": params["num_inference_steps"],
"guidance_scale": params["guidance_scale"],
"guidance_scale_2": params["guidance_scale_2"],
"embedded_cfg_scale": params["embedded_cfg_scale"],
"seed": params["seed"],
"fps": params["fps"],
"negative_prompt": params["negative_prompt"],
"c2ws_plucker_emb": c2ws_plucker_emb,
}
generation_kwargs = {
"output_path": output_dir,
"image_path": params["image_path"],
"height": params["height"],
"width": params["width"],
"num_frames": aligned_num_frames,
"num_inference_steps": params["num_inference_steps"],
"guidance_scale": params["guidance_scale"],
"guidance_scale_2": params["guidance_scale_2"],
"embedded_cfg_scale": params["embedded_cfg_scale"],
"seed": params["seed"],
"fps": params["fps"],
"negative_prompt": params["negative_prompt"],
"c2ws_plucker_emb": c2ws_plucker_emb,
}
generator: VideoGenerator | None = None
try:
generator = VideoGenerator.from_pretrained(
model_path=params["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
generator = VideoGenerator.from_pretrained(
model_path=params["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
finally:
if generator is not None:
generator.shutdown()
@@ -205,12 +205,12 @@ def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
assert os.path.exists(generated_video_path), (
f"Output video was not generated at {generated_video_path}")
reference_folder = build_reference_folder_path(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
reference_folder = build_reference_folder_path(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
if not os.path.exists(reference_folder):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
@@ -234,11 +234,11 @@ def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
mean_ssim = ssim_values[0]
logger.info("SSIM mean value: %s", mean_ssim)
write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path,
params["num_inference_steps"], prompt)
write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path,
params["num_inference_steps"], prompt)
min_acceptable_ssim = 0.90
min_acceptable_ssim = 0.70
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}")
@@ -0,0 +1,93 @@
# SPDX-License-Identifier: Apache-2.0
"""SSIM-based similarity test for LTX-2 distilled text-to-video.
Parameters derived from examples/inference/basic/basic_ltx2_distilled.py,
with resolution + num_inference_steps reduced to keep GPU CI runtime
bounded. Full-quality variant (via ``--ssim-full-quality``) falls back
to the ``ltx2_distilled`` preset defaults.
"""
import os
import pytest
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.logger import init_logger
from fastvideo.tests.ssim.inference_similarity_utils import (
resolve_inference_device_reference_folder,
run_text_to_video_similarity_test,
)
logger = init_logger(__name__)
REQUIRED_GPUS = 2
device_reference_folder = resolve_inference_device_reference_folder(logger)
LTX2_DISTILLED_PARAMS = {
"num_gpus": 2,
"model_path": "FastVideo/LTX2-Distilled-Diffusers",
"height": 512,
"width": 768,
"num_frames": 45,
"num_inference_steps": 4,
"guidance_scale": 1.0,
"seed": 10,
"sp_size": 2,
"tp_size": 1,
"fps": 24,
"ltx2_vae_tiling": True,
}
_LTX2_DISTILLED_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
LTX2_DISTILLED_PARAMS["model_path"])
LTX2_DISTILLED_FULL_QUALITY_PARAMS = {
"num_gpus": LTX2_DISTILLED_PARAMS["num_gpus"],
"model_path": LTX2_DISTILLED_PARAMS["model_path"],
"height": _LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.height,
"width": _LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.width,
"num_frames": _LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.num_frames,
"num_inference_steps":
_LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.num_inference_steps,
"guidance_scale": _LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.guidance_scale,
"seed": _LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.seed,
"sp_size": LTX2_DISTILLED_PARAMS["sp_size"],
"tp_size": LTX2_DISTILLED_PARAMS["tp_size"],
"fps": _LTX2_DISTILLED_FULL_QUALITY_DEFAULTS.fps,
"ltx2_vae_tiling": LTX2_DISTILLED_PARAMS["ltx2_vae_tiling"],
}
LTX2_DISTILLED_MODEL_TO_PARAMS = {
"LTX2-Distilled-Diffusers": LTX2_DISTILLED_PARAMS,
}
FULL_QUALITY_LTX2_DISTILLED_MODEL_TO_PARAMS = {
"LTX2-Distilled-Diffusers": LTX2_DISTILLED_FULL_QUALITY_PARAMS,
}
LTX2_DISTILLED_TEST_PROMPTS = [
"A warm sunny backyard. The camera starts in a tight cinematic "
"close-up of a woman and a man in their 30s, facing each other with "
"serious expressions. The camera slowly pans right, revealing a "
"grandfather in the garden wearing enormous butterfly wings, waving "
"his arms in the air like he's trying to take off. The tone is "
"deadpan, absurd, and quietly tragic.",
]
@pytest.mark.parametrize("prompt", LTX2_DISTILLED_TEST_PROMPTS)
@pytest.mark.parametrize("attention_backend_name", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(LTX2_DISTILLED_MODEL_TO_PARAMS.keys()))
def test_ltx2_distilled_inference_similarity(
prompt: str,
attention_backend_name: str,
model_id: str,
) -> None:
run_text_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=LTX2_DISTILLED_MODEL_TO_PARAMS,
full_quality_params_map=FULL_QUALITY_LTX2_DISTILLED_MODEL_TO_PARAMS,
min_acceptable_ssim=0.60,
)
@@ -145,6 +145,7 @@ TURBODIFFUSION_I2V_IMAGE_PATHS = [
]
@pytest.mark.skip(reason="Disabled: causes OOM too often in CI")
@pytest.mark.parametrize("prompt", TURBODIFFUSION_I2V_TEST_PROMPTS)
@pytest.mark.parametrize(
"model_id",
@@ -0,0 +1,91 @@
import numpy as np
import PIL.Image
import pytest
import torch
from fastvideo.pipelines.stages.image_encoding import ImageVAEEncodingStage
def make_stage() -> ImageVAEEncodingStage:
# Bypass __init__: preprocess() does not use self.vae.
return ImageVAEEncodingStage.__new__(ImageVAEEncodingStage)
def test_preprocess_pil_image():
stage = make_stage()
arr = np.array(
[[[0, 0, 0], [128, 128, 128], [255, 255, 255]]],
dtype=np.uint8,
)
image = PIL.Image.fromarray(arr, mode="RGB")
out = stage.preprocess(image, vae_scale_factor=1, height=1, width=3)
assert out.dtype == torch.float32
assert out.shape == (1, 3, 1, 3)
torch.testing.assert_close(
out[0, 0, 0],
torch.tensor([-1.0, 128.0 / 255.0 * 2 - 1, 1.0]),
atol=1e-6,
rtol=0,
)
def test_preprocess_uint8_tensor():
stage = make_stage()
image = torch.tensor(
[[[[0, 128, 255]], [[0, 128, 255]], [[0, 128, 255]]]],
dtype=torch.uint8,
)
out = stage.preprocess(image, vae_scale_factor=1, height=1, width=3)
assert out.dtype == torch.float32
expected = torch.tensor([-1.0, 128.0 / 255.0 * 2 - 1, 1.0])
torch.testing.assert_close(out[0, 0, 0], expected, atol=1e-6, rtol=0)
assert out.max().item() <= 1.0
assert out.min().item() >= -1.0
def test_preprocess_float01_tensor_matches_uint8_path():
stage = make_stage()
uint8_image = torch.tensor(
[[[[0, 128, 255]], [[0, 128, 255]], [[0, 128, 255]]]],
dtype=torch.uint8,
)
float_image = uint8_image.float() / 255.0
out_uint8 = stage.preprocess(uint8_image, vae_scale_factor=1, height=1, width=3)
out_float = stage.preprocess(float_image, vae_scale_factor=1, height=1, width=3)
torch.testing.assert_close(out_uint8, out_float, atol=1e-6, rtol=0)
def test_preprocess_already_normalized_passthrough():
stage = make_stage()
# Already in [-1, 1]; do_normalize branch must be skipped.
image = torch.tensor(
[[[[-1.0, 0.0, 1.0]], [[-1.0, 0.0, 1.0]], [[-1.0, 0.0, 1.0]]]],
dtype=torch.float32,
)
out = stage.preprocess(image, vae_scale_factor=1, height=1, width=3)
torch.testing.assert_close(out, image, atol=0, rtol=0)
@pytest.mark.parametrize(
"bad_input, expected_exc",
[
# Float tensor outside [-1, 1] / [0, 1].
(torch.tensor([[[[0.0, 1.5]]]], dtype=torch.float32), ValueError),
# Non-floating, non-uint8 tensor.
(torch.tensor([[[[0, 1]]]], dtype=torch.int32), ValueError),
# Wrong outer type.
(np.zeros((1, 3, 1, 3), dtype=np.float32), TypeError),
],
)
def test_preprocess_rejects_invalid_inputs(bad_input, expected_exc):
stage = make_stage()
with pytest.raises(expected_exc):
stage.preprocess(bad_input, vae_scale_factor=1, height=1, width=2)
@@ -666,6 +666,10 @@ class DistillationPipeline(TrainingPipeline):
scheduler=self.noise_scheduler).unflatten(
0, real_score_pred_noise_uncond.shape[:2])
# CFG on the real-score teacher. Uses the DMD2 parameterization
# x_cond + w * (x_cond - x_uncond), which is offset by 1 from the
# Ho & Salimans form x_uncond + w * (x_cond - x_uncond):
# w=0 -> cond, w=-1 -> uncond, w_standard = w + 1.
real_score_pred_video = pred_real_video_cond + (pred_real_video_cond -
pred_real_video_uncond) * self.real_score_guidance_scale
@@ -294,3 +294,52 @@ def test_ltx2_pipeline_smoke():
assert ref_video.shape == fastvideo_out.shape
assert_close(ref_video, fastvideo_out, atol=2 / 255, rtol=1e-3)
def test_ltx2_typed_surface_preflight() -> None:
"""Preflight: the PR 6 typed LTX-2 surface (preset + refine
override dataclasses + colocated pipeline config) must be importable
and registered before any GPU pipeline construction is attempted.
Pure-Python; does not need CUDA or model weights. Catches import-
wiring regressions (registry loss, renamed modules, preset dropped
from ALL_PRESETS) that would otherwise only surface on a GPU host.
"""
import fastvideo.registry # noqa: F401 — triggers preset registration
from fastvideo.api.presets import get_preset, get_presets_for_family
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
LTX2RefinePresetOverride,
LTX2RefineStageOverride,
refine_preset_override_fields,
refine_stage_override_fields,
)
from fastvideo.pipelines.basic.ltx2.stages import ( # noqa: F401
LTX2AudioDecodingStage,
LTX2DenoisingStage,
LTX2LatentPreparationStage,
LTX2TextEncodingStage,
)
# All three LTX-2 presets registered.
names = {p.name for p in get_presets_for_family("ltx2")}
assert names == {"ltx2_base", "ltx2_distilled", "ltx2_two_stage"}
# Two-stage preset has the denoise + refine topology and pulls its
# refine allowed_overrides from the typed dataclass.
two_stage = get_preset("ltx2_two_stage", "ltx2")
stage_names = [s.name for s in two_stage.stage_schemas]
assert stage_names == ["denoise", "refine"]
refine_spec = two_stage.stage_schemas[1]
assert refine_spec.allowed_overrides == refine_stage_override_fields()
# Override dataclasses are constructable and advertise disjoint
# field sets (init-time vs. per-request).
assert LTX2RefinePresetOverride().enabled is None
assert LTX2RefineStageOverride().num_inference_steps is None
preset_fields = refine_preset_override_fields()
stage_fields = refine_stage_override_fields()
assert preset_fields.isdisjoint(stage_fields)
# Colocated pipeline config is discoverable.
assert LTX2T2VConfig().vae_tiling is True
+1 -1
View File
@@ -12,7 +12,7 @@ from fastvideo.registry import (
get_pipeline_config_cls_from_name,
get_sampling_param_cls_for_name,
)
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
@pytest.mark.parametrize(