Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
49a13e6359 | ||
|
|
7b872cc41e | ||
|
|
37418946c8 | ||
|
|
95fd29e0cb | ||
|
|
e17cd2633c | ||
|
|
e0dc5f2b0c | ||
|
|
70ee5d230c | ||
|
|
24ced500f5 |
Executable
+96
@@ -0,0 +1,96 @@
|
||||
#!/usr/bin/env bash
|
||||
# Sync .agents/skills/ into .claude/skills/ via per-skill symlinks.
|
||||
#
|
||||
# Why: Claude Code only scans .claude/skills/ and ~/.claude/skills/ for
|
||||
# user-invocable skills (no skillsPath config exists — see
|
||||
# https://code.claude.com/docs/en/skills.md). This repo's skills live
|
||||
# in .agents/skills/ so they travel with the repo and stay under git.
|
||||
# Run this once after cloning (or after adding/removing a skill) to
|
||||
# expose them to Claude Code without maintaining a parallel tree.
|
||||
#
|
||||
# Usage:
|
||||
# .agents/scripts/sync-skills.sh
|
||||
#
|
||||
# Idempotent and safe to re-run. Prunes stale symlinks whose source
|
||||
# has been removed from .agents/skills/. Leaves hand-written
|
||||
# .claude/skills/<name>/ directories untouched (only symlinks are
|
||||
# managed).
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(git -C "$(dirname "$0")" rev-parse --show-toplevel)"
|
||||
SRC_DIR="$REPO_ROOT/.agents/skills"
|
||||
DST_DIR="$REPO_ROOT/.claude/skills"
|
||||
|
||||
if [[ ! -d "$SRC_DIR" ]]; then
|
||||
echo "Error: $SRC_DIR does not exist." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir -p "$DST_DIR"
|
||||
|
||||
linked=0
|
||||
unchanged=0
|
||||
skipped=0
|
||||
pruned=0
|
||||
|
||||
link_skill() {
|
||||
local name="$1"
|
||||
local src="$SRC_DIR/$name"
|
||||
local dst="$DST_DIR/$name"
|
||||
# Relative target keeps symlinks portable across clones.
|
||||
local rel="../../.agents/skills/$name"
|
||||
|
||||
if [[ -L "$dst" ]]; then
|
||||
if [[ "$(readlink "$dst")" == "$rel" ]]; then
|
||||
unchanged=$((unchanged + 1))
|
||||
return
|
||||
fi
|
||||
rm "$dst"
|
||||
elif [[ -e "$dst" ]]; then
|
||||
echo "Skipped (not a symlink): .claude/skills/$name" >&2
|
||||
skipped=$((skipped + 1))
|
||||
return
|
||||
fi
|
||||
|
||||
ln -s "$rel" "$dst"
|
||||
echo "Linked: .claude/skills/$name -> $rel"
|
||||
linked=$((linked + 1))
|
||||
}
|
||||
|
||||
prune_stale() {
|
||||
local link="$1"
|
||||
local target
|
||||
target="$(readlink "$link")"
|
||||
case "$target" in
|
||||
../../.agents/skills/*) ;;
|
||||
*) return ;;
|
||||
esac
|
||||
local name="${target##*/}"
|
||||
if [[ ! -d "$SRC_DIR/$name" ]]; then
|
||||
rm "$link"
|
||||
echo "Pruned stale: .claude/skills/$(basename "$link")"
|
||||
pruned=$((pruned + 1))
|
||||
fi
|
||||
}
|
||||
|
||||
for src in "$SRC_DIR"/*/; do
|
||||
[[ -d "$src" ]] || continue
|
||||
name="$(basename "$src")"
|
||||
# Only treat directories that actually contain a SKILL.md as skills.
|
||||
[[ -f "$src/SKILL.md" ]] || continue
|
||||
link_skill "$name"
|
||||
done
|
||||
|
||||
shopt -s nullglob
|
||||
for link in "$DST_DIR"/*; do
|
||||
[[ -L "$link" ]] || continue
|
||||
prune_stale "$link"
|
||||
done
|
||||
shopt -u nullglob
|
||||
|
||||
printf "\nSummary: %d linked, %d unchanged, %d pruned" "$linked" "$unchanged" "$pruned"
|
||||
if [[ "$skipped" -gt 0 ]]; then
|
||||
printf ", %d skipped (non-symlink collision)" "$skipped"
|
||||
fi
|
||||
printf "\n"
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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.
|
||||
@@ -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"}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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`.
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -11,10 +11,34 @@ import torch
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5ArchConfig(T5ArchConfig):
|
||||
"""T5 arch that pads tokenizer output to ``max_length``.
|
||||
|
||||
LongCat's denoising stage concatenates positive and negative
|
||||
attention masks along the batch dimension for CFG, which requires
|
||||
uniform seq length. The shared :class:`T5ArchConfig` dropped the
|
||||
``"padding": "max_length"`` tokenizer kwarg so other DiTs could run
|
||||
with variable-length masks; LongCat still needs the uniform
|
||||
contract.
|
||||
"""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5Config(T5Config):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LongCatT5ArchConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields."""
|
||||
@@ -103,8 +127,9 @@ class LongCatT2V480PConfig(PipelineConfig):
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
# UMT5 uses T5-like config; postprocess pads to 512. LongCatT5Config
|
||||
# restores ``padding="max_length"`` for the CFG concat contract.
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (LongCatT5Config(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (longcat_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (umt5_postprocess_text, ))
|
||||
|
||||
@@ -1,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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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}})
|
||||
@@ -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": {},
|
||||
},
|
||||
|
||||
@@ -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())
|
||||
@@ -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"),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user