Compare commits

..
Author SHA1 Message Date
SolitaryThinker 4a427692bf update 2026-04-21 10:26:57 -07:00
SolitaryThinker 1214fb0f74 docs(seed-ssim-skill): fix modal volume get path + force flag
Post-first-run corrections to the seed-ssim-references skill:
- modal volume get needs --force when the local parent directory
  already exists, otherwise it errors with [Errno 21] Is a directory.
- The downloaded tree has an extra generated_videos/ level (from the
  volume layout in _sync_generated_videos_to_volume), so copy-local's
  --generated-dir must include it.
2026-04-21 10:22:36 -07:00
SolitaryThinker 2dddbdf4c0 fix(kernel-build): detect CUDA arch via active venv python directly
uv run --active --no-project can provision its own interpreter on
some uv versions and misses packages installed into VIRTUAL_ENV,
causing ModuleNotFoundError: torch right after uv pip install -e
.[test] succeeded. Invoke $VIRTUAL_ENV/bin/python directly so
detect_with_torch reliably sees the freshly-installed torch.
2026-04-21 10:06:07 -07:00
SolitaryThinker 4aa065b96c fix(ssim-modal): install torch before fastvideo-kernel build
fastvideo-kernel/build.sh needs torch to detect the host CUDA arch via
detect_with_torch. The previous order built the kernel before uv pip
install -e .[test], so torch wasn't present yet and detection failed
with ModuleNotFoundError. Install the package (pulling torch) first,
then build the kernel from source to shadow the PyPI wheel.
2026-04-21 09:58:19 -07:00
SolitaryThinker 7b1cd12059 update 2026-04-21 09:44:00 -07:00
SolitaryThinkerandClaude Opus 4.7 d041b038bf [misc] [6/n] Improve API: tidy LTX-2 refine override helpers
Promote the refine-override field accessors to module-level frozenset
constants (``REFINE_PRESET_OVERRIDE_FIELDS``,
``REFINE_STAGE_OVERRIDE_FIELDS``, ``REFINE_FLAT_KEYS``) so callers
reference a single source of truth instead of recomputing on each
call. Drop the redundant ``dict(deepcopy(value))`` double-copy on the
``torch_compile_kwargs`` legacy path, compress the
``_compile_config_to_torch_kwargs`` docstring, and remove a handful of
comments that just restated the code or named a future PR. Narrow the
schema-parity ``walk_packages`` traversal to
``configs/pipelines/*`` and ``basic/<family>/pipeline_configs`` so the
test no longer imports every heavy model module under ``basic/``.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-21 08:32:19 -07:00
SolitaryThinker 9a15721d28 [misc] [6/n] Simplify LTX-2 preset scaffolding
Post-review cleanup (no behavior change):

Correctness:
  * compat.py refine reverse-flatten now iterates
    _LTX2_REFINE_FLAT_KEYS, which unions the typed
    LTX2RefinePresetOverride and LTX2RefineStageOverride field sets.
    The previous hardcoded 4-tuple missed image_crf and
    video_position_offset_sec, so any init-time
    preset_overrides.refine.image_crf would silently drop on reverse.
    New regression test TestRefineFlattenCoversAllTypedFields pins the
    invariant.

Simplification:
  * refine_preset_override_to_dict and refine_stage_override_to_dict
    were byte-identical; collapse to a single refine_override_to_dict.
  * Trim narrative / PR-migration prose from docstrings and re-export
    shims: stage_overrides.py module + field docstrings, schema.py
    CompileConfig + PipelineSelection.vae_tiling, presets.py
    LTX2_TWO_STAGE header, configs/pipelines/__init__.py and
    pipelines/stages/__init__.py and ltx2/stages/__init__.py.
2026-04-21 08:32:19 -07:00
SolitaryThinker 185d833d68 [test] [6/n] Improve API: gpu_pool-style LTX-2 kwarg round-trip
Close out PR 6 with the compat mappings for every flat LTX-2 kwarg
the FastVideo-internal ui/ltx2-streaming/server/gpu_pool.py passes to
VideoGenerator.from_pretrained, plus an integration test that freezes
the gpu_pool load_kwargs dict as a parity guard.

Compat additions (fastvideo/api/compat.py):

  legacy_from_pretrained_to_config forward routing for:
    - config_model_path -> components.config_root
    - ltx2_refine_enabled -> preset_overrides.refine.enabled
    - ltx2_refine_upsampler_path -> components.upsampler_weights
      (empty string collapses to None)
    - ltx2_refine_lora_path -> components.lora_path
      (empty string collapses to None)
    - ltx2_refine_add_noise -> preset_overrides.refine.add_noise
    - ltx2_refine_num_inference_steps -> preset_overrides.refine.num_inference_steps
    - ltx2_refine_guidance_scale -> preset_overrides.refine.guidance_scale

  generator_config_to_fastvideo_args reverse routing:
    - components.config_root -> config_model_path
    - components.upsampler_weights -> ltx2_refine_upsampler_path
      (no longer raises NotImplementedError)
    - preset_overrides.refine.{...} now flattens back to
      ltx2_refine_{enabled, add_noise, num_inference_steps, guidance_scale}
      instead of leaking as a nested "refine" kwarg.

PR 7.6 (gpu_pool upstream) and the Dynamo adapter can now build a typed
GeneratorConfig from their CLI without knowing any legacy LTX-2 kwarg
name — the hard gate called out in PR plan.md § "Why the gpu_pool.py
typed-replacement scope matters".

Tests (fastvideo/tests/api/test_ltx2_gpu_pool_translation.py):

  - GPU_POOL_LOAD_KWARGS fixture mirrors gpu_pool.py lines 233-260
    (minus opaque objects: PipelineConfig instance, the text-encoder
    torch.compile flag not yet in the public FastVideoArgs).
  - TestGpuPoolForwardTranslation asserts each flat kwarg lands on the
    correct typed field and that pipeline.experimental stays empty
    (no silent fallthrough).
  - TestGpuPoolReverseTranslation asserts the same dict reproduces on
    the FastVideoArgs.from_kwargs call — including the four ltx2_refine_*
    fields, ltx2_vae_tiling, torch_compile_kwargs, and config_model_path.
  - TestCompileExtrasPreserved checks non-typed torch.compile kwargs
    (options, disable) ride through CompileConfig.extras in both
    directions.

All 148 API tests pass locally (227 including entrypoints).
2026-04-21 08:32:19 -07:00
SolitaryThinker 5568c90591 [refactor] [6/n] Improve API: colocate LTX-2 pipeline config and stages
Move the LTX-2 family's PipelineConfig and the four LTX-2-specific
stage modules into fastvideo/pipelines/basic/ltx2/ so each model family
directory is self-contained (pipeline implementation + presets +
stage_overrides + pipeline_configs + stages in one place). See the
"Pipeline Package Structure" section in PR plan.md and the per-model
colocation plan in apirefactor.md Phase 8.

File moves (git mv preserves history):

  fastvideo/configs/pipelines/ltx2.py
    -> fastvideo/pipelines/basic/ltx2/pipeline_configs.py
  fastvideo/pipelines/stages/ltx2_audio_decoding.py
    -> fastvideo/pipelines/basic/ltx2/stages/ltx2_audio_decoding.py
  fastvideo/pipelines/stages/ltx2_denoising.py
    -> fastvideo/pipelines/basic/ltx2/stages/ltx2_denoising.py
  fastvideo/pipelines/stages/ltx2_latent_preparation.py
    -> fastvideo/pipelines/basic/ltx2/stages/ltx2_latent_preparation.py
  fastvideo/pipelines/stages/ltx2_text_encoding.py
    -> fastvideo/pipelines/basic/ltx2/stages/ltx2_text_encoding.py

Import/export updates:

  - fastvideo/registry.py and tests/local_tests/test_ltx2_registry.py
    now import LTX2T2VConfig from the new colocated path directly.
  - fastvideo/configs/pipelines/__init__.py keeps the legacy
    `fastvideo.configs.pipelines.LTX2T2VConfig` export by re-exporting
    from the new location — external users of the old path keep working.
  - fastvideo/pipelines/stages/__init__.py does the same for the four
    LTX-2 stage classes, re-exporting from the new
    pipelines/basic/ltx2/stages/ subpackage.
  - new fastvideo/pipelines/basic/ltx2/stages/__init__.py is the
    canonical home for the stage exports.

Parity-inventory walker (fastvideo/tests/api/test_schema_parity_inventory.py):
  _get_extra_dataclass_fields now accepts a tuple of package roots and
  descends recursively via pkgutil.walk_packages, so it discovers
  PipelineConfig subclasses that have moved out of
  fastvideo.configs.pipelines into fastvideo.pipelines.basic.<family>.
  The pipeline_config_extensions classification test now walks both
  roots so the LTX-2 fields (vocoder_config, audio_decoder_config,
  vocoder_precision, audio_decoder_precision) stay accounted for.

Nothing else changes: no behavior change; no test changes beyond the
walker; all 210 API+entrypoints tests still pass.
2026-04-21 08:32:19 -07:00
SolitaryThinker 8031f27817 [feat] [6/n] Improve API: typed torch.compile kwargs + pipeline.vae_tiling
Promote the four torch.compile kwargs the FastVideo-internal
ltx2-streaming gpu_pool hard-codes in its torch_compile_kwargs dict
into first-class typed fields on CompileConfig, and add a typed home
for ltx2_vae_tiling on PipelineSelection so PR 7.6's public gpu_pool
upstream has a typed boundary to hit.

Schema changes (fastvideo/api/schema.py):

  CompileConfig:
    - rename kwargs -> extras (breaking; pre-stable API)
    - add backend: str | None (e.g. "inductor")
    - add fullgraph: bool | None
    - add mode: str | None (e.g. "max-autotune-no-cudagraphs")
    - add dynamic: bool | None
    extras still holds uncommon kwargs (e.g. options, disable).

  PipelineSelection:
    - add vae_tiling: bool | None. Shared across model families that
      expose VAE tiling (LTX-2, Wan configs, etc.); None leaves the
      model default in place.

Compat (fastvideo/api/compat.py):

  - legacy_from_pretrained_to_config splits incoming
    torch_compile_kwargs={backend, fullgraph, mode, dynamic, ...}
    across the four typed fields and drops the remainder into extras.
  - generator_config_to_fastvideo_args reconstructs the flat
    torch_compile_kwargs dict via _compile_config_to_torch_kwargs,
    emitting only user-set typed fields plus extras.
  - legacy ltx2_vae_tiling routes to pipeline.vae_tiling on the way in
    and back to ltx2_vae_tiling on the way out.

Parity inventory (docs/design/inference_schema_parity_inventory.yaml):
  - torch_compile_kwargs now maps to the comma-separated leaf
    generator.engine.compile.backend,fullgraph,mode,dynamic,extras
  - ltx2_vae_tiling reclassified from preset_owned nested dict
    (.preset_overrides.ltx2.vae_tiling) to moved (.pipeline.vae_tiling)

11 new tests in fastvideo/tests/api/test_compat_translation.py cover
typed-key promotion, extras merge, None suppression, round-trip back
to FastVideoArgs, and the vae_tiling forward/reverse path. Parser
round-trip fixture updated for the new schema fields.
2026-04-21 08:32:19 -07:00
SolitaryThinker 0b7e8b5d1d [feat] [6/n] Improve API: typed LTX2 refine stage overrides
Add typed public dataclasses that describe the LTX-2 refine override
surfaces, and bind them to the ltx2_two_stage preset's stage schema so
the validation-layer field set stays in lockstep with the dataclass:

  * LTX2RefinePresetOverride (init-time, for preset_overrides.refine):
    - enabled: bool | None  — toggle the refine stage topology
    - add_noise: bool | None — controls LTX2UpsampleStage noise mixing

  * LTX2RefineStageOverride (per-request, for stage_overrides.refine):
    - num_inference_steps: int | None — stage-2 denoise steps (2 or 3)
    - guidance_scale: float | None — force_guidance_scale for stage-2
    - image_crf: int | None — image-encoding CRF hint
    - video_position_offset_sec: float | None — RoPE shift for audio
      conditioning continuation

Asset wiring (upsampler weights, refine LoRA) stays on
ComponentConfig.upsampler_weights / .lora_path — no new typed home is
added here because those already exist.

The ltx2_two_stage refine stage schema's allowed_overrides now reads
from refine_stage_override_fields() so the dataclass is the single
source of truth. Serialisation helpers
(refine_{preset,stage}_override_to_dict) drop None entries so only
user-set fields flow through the preset/stage override dicts.

The runtime still reads the refine knobs off FastVideoArgs at
pipeline-construction time; wiring these typed objects through compat
+ runtime is the next slice. This commit ships only the public
typed surface and the preset/dataclass invariant.

12 new tests in fastvideo/tests/api/test_ltx2_stage_overrides.py cover
default None construction, explicit construction, to_dict() dropping
None, fields() accessor, and an invariant test proving the preset's
refine stage allowed_overrides exactly matches the dataclass fields.
2026-04-21 08:32:19 -07:00
SolitaryThinker 279e52ad8d [feat] [6/n] Improve API: add ltx2_two_stage preset
Land the ltx2_two_stage inference preset (PR 6 commit 1/5) that exposes
the public LTX-2 two-stage distilled flow: half-resolution denoise
followed by 2x spatial upsample + stage-2 refine denoise (3 steps with
the official distilled sigma schedule, or 2 steps with the reduced
variant).

Schema:
  - new _REFINE_STAGE PresetStageSpec with kind="refinement" and
    allowed per-request overrides {num_inference_steps, guidance_scale,
    image_crf, video_position_offset_sec}
  - new LTX2_TWO_STAGE preset with stage_schemas=(denoise, refine) and
    stage_defaults.refine={num_inference_steps: 2, guidance_scale: 1.0}
    matching the load_kwargs the internal ltx2-streaming gpu_pool uses
  - LTX2_TWO_STAGE registered via ALL_PRESETS -> _register_presets()

Init-time refine wiring (enabled, add_noise, upsampler/LoRA/transformer
paths) flows through generator.pipeline.preset_overrides.refine.* and
generator.pipeline.components.{upsampler_weights, lora_path} in a
follow-up commit — the refine stage's internal implementation
(fastvideo/pipelines/stages/ltx2_refine.py in FastVideo-internal) already
treats those as FastVideoArgs fields; this commit only ships the public
preset surface.

5 new tests under fastvideo/tests/api/test_presets.py::TestLtx2Presets:
registration count (3 presets now), two-stage topology, stage defaults,
valid per-request refine override keys, and unknown-override rejection.
2026-04-21 08:32:19 -07:00
SolitaryThinker 6b3c1223c6 [misc] add .agents/scripts/sync-skills.sh
Claude Code only scans ~/.claude/skills/ and .claude/skills/ for
user-invocable skills (no skillsPath / skillsDir config option
exists — https://code.claude.com/docs/en/skills.md). Skills in this
repo live under .agents/skills/ so they travel with the repo and
stay under git.

Add a one-shot idempotent sync script that symlinks each
.agents/skills/<name>/ directory into .claude/skills/<name>. Run
after cloning or after adding/removing a skill:

    .agents/scripts/sync-skills.sh

Behavior:
  * relative symlinks (../../.agents/skills/<name>) so the link
    survives moving the clone
  * requires a SKILL.md inside each skill directory to be eligible
  * prunes stale symlinks whose source vanished from .agents/skills/
  * refuses to clobber a pre-existing .claude/skills/<name>/ that
    isn't a symlink (lets the operator keep hand-written skills
    alongside managed ones)
  * prints a linked / unchanged / pruned / skipped summary

Note: .claude/ is gitignored; the symlinks themselves are not
committed. The script is the source of truth for reconstructing
them.
2026-04-17 18:26:06 -07:00
SolitaryThinker 0c8687c919 [misc] add seed-ssim-references agent skill
Wraps the existing fastvideo/tests/modal/ssim_test.py
--sync-generated-to-volume path with a step-by-step SKILL.md and a
thin seed_ssim.sh launcher so new SSIM tests can bootstrap their
HF reference videos without manual Modal wrangling.

Motivation: the new test_ltx2_similarity.py (committed earlier on
this branch) has no HF reference video yet; the next operator needs
a deterministic procedure for generating + uploading the first set.
Generalises beyond LTX-2 — any new family test can reuse verbatim.
2026-04-17 17:59:21 -07:00
SolitaryThinker ddbf41fa0e [test] add LTX-2 distilled T2V SSIM regression test
LTX-2 was the only model family in fastvideo/pipelines/basic/ without an
SSIM coverage file. Add test_ltx2_similarity.py alongside the other
per-family SSIM tests so the in-flight API refactor and future
LTX-2-specific changes (two-stage refine, gpu_pool upstream, Dynamo
backend) have a golden-quality regression guard.

Parameters:

  Default (CI-friendly):
    model: FastVideo/LTX2-Distilled-Diffusers (8-step distilled)
    resolution: 512x768, num_frames=45, num_inference_steps=4
    sp_size=2 on 2 GPUs, FLASH_ATTN backend
    ltx2_vae_tiling=True for peak-memory safety

  --ssim-full-quality:
    falls back to the ltx2_distilled preset defaults
    (1024x1536, 121 frames, 8 steps, guidance_scale=1.0)

Uses the shared run_text_to_video_similarity_test helper (same pattern
as Wan/TurboDiffusion), so _build_init_kwargs picks up
ltx2_vae_tiling + related tile sizes automatically.

REQUIRED_GPUS = 2 is declared at module scope so the Modal SSIM
orchestrator (fastvideo/tests/modal/ssim_test.py) schedules it
correctly on the L40S:8 runner. LTX2_DISTILLED_MODEL_TO_PARAMS is
named so the orchestrator can split by model id (one subprocess per
entry) for future multi-model LTX-2 coverage.

Threshold min_acceptable_ssim=0.93 matches Wan T2V.

Reference videos are not in the repo. After this lands on main, run
the test once on an L40S and upload via
`python fastvideo/tests/ssim/reference_videos_cli.py upload --quality-tier all`
to seed FastVideo/ssim-reference-videos. Subsequent runs (including
regression guards for in-flight refactor PRs) auto-download the
references before the test executes.
2026-04-17 16:35:51 -07:00
53 changed files with 267 additions and 8113 deletions
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
-230
View File
@@ -1,230 +0,0 @@
# 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.
-1
View File
@@ -6,4 +6,3 @@
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
{"name": "review-add-model-pr", "description": "Review a PR that adds a new model / pipeline (or major variant) to FastVideo. Walks the canonical surface, indexes the 36 documented add-model failure modes by where they show up in a diff, and produces a structured verdict (block / nit / follow-up).", "path": "review-add-model-pr/SKILL.md", "status": "draft", "trust": "low"}
-482
View File
@@ -1,482 +0,0 @@
---
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).
@@ -354,8 +354,6 @@ surfaces:
return_frames: request.output.return_frames
return_trajectory_latents: request.runtime.return_trajectory_latents
return_trajectory_decoded: request.runtime.return_trajectory_decoded
continuation_state: request.state
return_continuation_state: request.output.return_state
preset_owned:
t_thresh: request.stage_overrides.refine.t_thresh
spatial_refine_only: request.stage_overrides.refine.spatial_refine_only
-177
View File
@@ -1,177 +0,0 @@
# Streaming WebSocket Server Contract
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
document is the authoritative spec for the message catalogue and the
session state machine. Any change to either must update this document
in the same PR that touches `protocol.py` or `session.py`.
## Endpoint
| Path | Protocol | Purpose |
|---|---|---|
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
The server is launched by `fastvideo serve --config <serve.yaml>` when
the config carries a `streaming:` block. Without that block the same CLI
launches the OpenAI stateless HTTP server instead.
## Connection lifecycle
Every WebSocket connection holds exactly one `Session`. Sessions move
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
```
┌──────────────┐
│ INITIALIZING │ ← WebSocket accepted, before init frame
└──────┬───────┘
│ session_init_v2 received
┌──────────────┼──────────────┐
▼ ▼ ▼
QUEUED GPU_BINDING REJECTED
│ │ ↑
│ slot ready │ │ max-sessions hit
▼ ▼ │ or invalid init
┌────────┐ │
│ ACTIVE │ ────────┘
└────┬───┘
segment loop │
│
┌───────────┼───────────┐
▼ ▼ ▼
COMPLETE ERROR TIMEOUT
(clean leave) (any failure) (idle / segment_cap reached)
```
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
no transitions out. The transition matrix is enforced in
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
`SessionManager` enforces the per-process budgets pulled from
`StreamingConfig`:
- `session_timeout_seconds` — idle reaper drops sessions that haven't
advanced; non-terminal sessions transition to `TIMEOUT`.
- `generation_segment_cap` — a session that hits the cap transitions to
`COMPLETE` after the last segment ships.
## Message catalogue
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
`protocol.py` are the source of truth; this table is the human-readable
view.
### Client → server
| `type` | Required fields | Purpose |
|---|---|---|
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
The opening frame must be `session_init_v2`. Any other first frame is
rejected with an `error` (code `invalid_message`) and the WebSocket is
closed.
### Server → client
| `type` | Carries | When emitted |
|---|---|---|
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
## Continuation state
The session optionally accepts a `continuation_state` dict inside the
opening `session_init_v2` frame. When present, the server hydrates it
into a `ContinuationState(kind, payload)` envelope and feeds it as the
`request.state` on the first segment's `GenerationRequest` — letting a
client resume after a disconnect, migrate sessions across processes,
or replay a prior session.
After every segment, if the runtime returns a fresh state, the server
persists it to the `SessionStore` so a `snapshot_state` request can
export it. The store and serialization contracts live with the model
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
## Example flow
```
client server
────── ──────
WS /v1/stream ─────── connect ─────────────────────────►
◄────── (accept)
{"type": "session_init_v2",
"preset": "ltx2_two_stage",
"curated_prompts": ["a fox in snow", "the fox jumps"],
"initial_image": {...},
"stream_mode": "av_fmp4"} ─────────────────────────────►
(validate, queue, bind)
◄──── {"type": "queue_status",
"position": 0, "queue_depth": 0}
◄──── {"type": "gpu_assigned",
"gpu_id": 0, "model_id": "..."}
◄──── {"type": "ltx2_stream_start", ...}
{"type": "segment_prompt_source",
"prompt": "a fox in snow",
"source": "curated"} ───────────────────────────────────►
(run pipeline)
◄──── {"type": "ltx2_segment_start",
"segment_idx": 1, ...}
◄──── {"type": "step_complete",
"segment_idx": 1, "timings": {...}}
◄──── {"type": "media_init",
"segment_idx": 1,
"mime": "video/mp4", ...}
◄──── <binary fMP4 init segment>
◄──── <binary fMP4 fragment>
◄──── <binary fMP4 fragment>
◄──── {"type": "media_segment_complete",
"segment_idx": 1, "chunks": 12}
◄──── {"type": "ltx2_segment_complete",
"segment_idx": 1, ...}
{"type": "segment_prompt_source",
"prompt": "the fox jumps"} ─────────────────────────────►
(segment 2 …)
{"type": "snapshot_state"} ──────────────────────────────►
◄──── {"type": "continuation_state_snapshot",
"kind": "ltx2.v1",
"payload": {"schema_version": 1, ...}}
(close) ──────────────────────────────────────────────────►
(session → COMPLETE)
```
## Backward / forward compatibility
- Adding a new client message: append a Pydantic model to `protocol.py`
with a unique `type`; add the discriminator entry to `ClientMessage`;
add a row to the table above. Old clients that don't send the new
message remain compatible.
- Adding a new server message: emit only when a new feature flag is
enabled (or always emit, since clients ignore unknown types).
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
→ `session_init_v3`) and accept both for one release cycle. Never
silently change field semantics under the same `type`.
-22
View File
@@ -86,25 +86,3 @@ 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.
@@ -5,11 +5,6 @@ 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,
@@ -27,8 +22,6 @@ 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,5 +1,3 @@
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
from __future__ import annotations
import os
@@ -8,11 +6,6 @@ from typing import Tuple
import torch
# ---------------------------------------------------------------------------
# Backend selection helpers
# ---------------------------------------------------------------------------
def _get_sm90_ops():
try:
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
@@ -32,66 +25,38 @@ 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"
# ---------------------------------------------------------------------------
# Index helpers
# ---------------------------------------------------------------------------
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."""
"""
Preferred map->index conversion used by the wrapper.
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
"""
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]), "
f"got shape={tuple(block_map.shape)}"
)
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), 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
except Exception as e: # pragma: no cover - environment issue
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
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=(),
@@ -101,40 +66,34 @@ def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
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
triton_block_sparse_attn_forward,
)
o, M = triton_block_sparse_attn_forward(
q.contiguous(),
k.contiguous(),
v.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
o, M = triton_block_sparse_attn_forward(q, k, v, 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,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: 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
@@ -150,32 +109,20 @@ def block_sparse_attn_backward_triton(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
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
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.contiguous(),
q.contiguous(),
k.contiguous(),
v.contiguous(),
o,
M,
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
variable_block_sizes,
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
)
return dq, dk, dv
@@ -188,8 +135,7 @@ def _block_sparse_attn_backward_triton_fake(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q)
@@ -198,28 +144,19 @@ def _block_sparse_attn_backward_triton_fake(
return dq, dk, dv
def _setup_context_triton(ctx, inputs, output):
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
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
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
block_sparse_attn_triton.register_autograd(
_backward_triton, setup_context=_setup_context_triton
)
def _setup_context_triton(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
# ---------------------------------------------------------------------------
# SM90 backend custom ops (index-native)
# ---------------------------------------------------------------------------
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
@torch.library.custom_op(
@@ -231,21 +168,21 @@ def block_sparse_attn_sm90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: 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.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
)
return o_padded, lse_padded
@@ -255,16 +192,11 @@ def _block_sparse_attn_sm90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: 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
@@ -280,34 +212,30 @@ def block_sparse_attn_backward_sm90(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: 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")
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
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())
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
dq, dk, dv = block_sparse_bwd(
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
q_padded,
k_padded,
v_padded,
o_padded,
lse_padded,
grad_output_padded.contiguous(),
grad_output_padded,
k2q_idx,
k2q_num,
variable_block_sizes,
variable_block_sizes.int(),
)
# 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)
# 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)
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
@@ -318,8 +246,7 @@ def _block_sparse_attn_backward_sm90_fake(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q_padded)
@@ -328,57 +255,21 @@ def _block_sparse_attn_backward_sm90_fake(
return dq, dk, dv
def _setup_context_sm90(ctx, inputs, output):
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, lse = output
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
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, q2k_idx, q2k_num, variable_block_sizes
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
)
return dq, dk, dv, None, None, None
return dq, dk, dv, None, None
block_sparse_attn_sm90.register_autograd(
_backward_sm90, setup_context=_setup_context_sm90
)
def _setup_context_sm90(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
o, lse = output
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
# ---------------------------------------------------------------------------
# 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)
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
def block_sparse_attn(
@@ -388,8 +279,16 @@ def block_sparse_attn(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""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
)
"""
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)
@@ -1,6 +1,6 @@
import math
import torch
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
from .block_sparse_attn import block_sparse_attn
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
# Try to load the C++ extension
@@ -125,18 +125,13 @@ def video_sparse_attn(
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, q_seq_len, dim)
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
# Sparse branch
topk_idx = torch.topk(scores, topk, dim=-1).indices
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]
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]
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
@@ -1,10 +1,9 @@
## 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,
@@ -154,114 +153,3 @@ 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
+8 -63
View File
@@ -17,7 +17,6 @@ from fastvideo.api.request_metadata import (
)
from fastvideo.api.schema import (
CompileConfig,
ContinuationState,
GenerationRequest,
GeneratorConfig,
InputConfig,
@@ -27,10 +26,7 @@ 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.pipelines.basic.ltx2.stage_overrides import REFINE_FLAT_KEYS
from fastvideo.utils import shallow_asdict
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
@@ -44,10 +40,7 @@ _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:
@@ -118,10 +111,8 @@ 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":
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
remaining: dict[str, Any] = (dict(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)
@@ -237,12 +228,6 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
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:
@@ -273,7 +258,7 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
refine = preset_overrides.pop("refine", None)
if isinstance(refine, Mapping):
for key in _LTX2_REFINE_FLAT_KEYS:
for key in REFINE_FLAT_KEYS:
if key in refine:
kwargs[f"ltx2_refine_{key}"] = refine[key]
kwargs.update(preset_overrides)
@@ -326,13 +311,10 @@ 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():
@@ -375,14 +357,9 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
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.
"""
"""Flatten typed ``CompileConfig`` back to the legacy
``torch_compile_kwargs`` dict, emitting only explicitly-set typed
fields and merging ``extras`` on top."""
out: dict[str, Any] = {}
for key in _COMPILE_TYPED_KEYS:
value = getattr(compile_config, key)
@@ -553,37 +530,6 @@ 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,
@@ -617,7 +563,6 @@ __all__ = [
"load_generator_config_from_file",
"normalize_generation_request",
"normalize_generator_config",
"register_continuation_kind",
"request_to_pipeline_overrides",
"request_to_sampling_param",
]
+6 -21
View File
@@ -1,16 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import copy
from dataclasses import dataclass, field, fields
from typing import TYPE_CHECKING, Any
from typing import 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__)
@@ -97,13 +92,9 @@ class SamplingParam:
movement_distance: float | None = None
camera_rotation: str | None = None
# LTX-2 multi-modal CFG and STG.
# cfg_scale defaults are 1.0 (CFG off) so ``ForwardBatch.__post_init__``
# doesn't force ``do_classifier_free_guidance`` on non-LTX-2 models that
# never override these fields. LTX-2 presets that need text-CFG on set
# them in their ``defaults`` dict (e.g. ``ltx2_base``).
ltx2_cfg_scale_video: float = 1.0
ltx2_cfg_scale_audio: float = 1.0
# LTX2 multi-modal CFG and STG
ltx2_cfg_scale_video: float = 3.0
ltx2_cfg_scale_audio: float = 7.0
ltx2_modality_scale_video: float = 3.0
ltx2_modality_scale_audio: float = 3.0
ltx2_rescale_scale: float = 0.7
@@ -112,12 +103,6 @@ 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
@@ -142,7 +127,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
@@ -158,7 +143,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
-6
View File
@@ -41,12 +41,6 @@ class CompileConfig:
"""
enabled: bool = False
text_encoder_enabled: bool | None = None
"""Whether ``torch.compile`` is applied to the text encoder. ``None``
keeps the runtime default. The public ``FastVideoArgs`` adapter does
not yet consume this flag; reserved so the realtime runtime upstream
(PR 7.6) has a typed home for its ``enable_torch_compile_text_encoder``
kwarg without routing through ``pipeline.experimental``."""
backend: str | None = None
fullgraph: bool | None = None
mode: str | None = None
+2 -27
View File
@@ -11,34 +11,10 @@ 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."""
@@ -127,9 +103,8 @@ class LongCatT2V480PConfig(PipelineConfig):
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
# 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(), ))
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
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, ))
+2 -29
View File
@@ -1,31 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.entrypoints.streaming.server import build_app, run_server
from fastvideo.entrypoints.streaming.session import (
Session,
SessionManager,
SessionState,
)
from fastvideo.entrypoints.streaming.session_store import (
BlobStore,
InMemoryBlobStore,
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
)
from fastvideo.entrypoints.streaming.server import run_server
__all__ = [
"BlobStore",
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
"InMemoryBlobStore",
"InMemorySessionStore",
"Session",
"SessionManager",
"SessionState",
"SessionStore",
"build_app",
"run_server",
]
__all__ = ["run_server"]
-252
View File
@@ -1,252 +0,0 @@
# 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",
]
+4 -520
View File
@@ -1,531 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-generator FastAPI + WebSocket streaming server."""
from __future__ import annotations
import asyncio
import contextlib
import os
import time
from dataclasses import dataclass
from typing import Any, Protocol
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
InputConfig,
OutputConfig,
SamplingConfig,
ServeConfig,
)
from fastvideo.entrypoints.streaming.protocol import (
AutoExtensionUpdated,
ContinuationStateSnapshot,
EnhancementUpdated,
ErrorMessage,
GenerationPausedUpdated,
GpuAssigned,
LoopGenerationUpdated,
Ltx2SegmentComplete,
Ltx2SegmentStart,
Ltx2StreamComplete,
Ltx2StreamStart,
MediaInit,
MediaSegmentComplete,
QueueStatus,
SeedPromptsUpdated,
SegmentPromptSource,
SessionInitV2,
SnapshotState,
StepComplete,
parse_client_message,
)
from fastvideo.entrypoints.streaming.session import (
InvalidSessionTransition,
Session,
SessionManager,
SessionRejected,
SessionState,
)
from fastvideo.entrypoints.streaming.session_init_image import (
persist_session_init_image, )
from fastvideo.entrypoints.streaming.session_store import (
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
from fastvideo.api.schema import ServeConfig
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# RFC 6455 WebSocket close codes used by the server.
_WS_CLOSE_UNSUPPORTED_DATA = 1003
_WS_CLOSE_TRY_AGAIN_LATER = 1013
class _GeneratorProto(Protocol):
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
def generate(self, request: GenerationRequest) -> Any:
...
@dataclass
class ServerState:
serve_config: ServeConfig
generator: _GeneratorProto
sessions: SessionManager
session_store: SessionStore
def build_app(
serve_config: ServeConfig,
generator: _GeneratorProto,
*,
session_store: SessionStore | None = None,
) -> FastAPI:
"""Build the FastAPI app used by :func:`run_server`.
Exposed so tests can drive the WebSocket endpoint in-process via
``starlette.testclient.TestClient(app).websocket_connect(...)``.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
"server; got None. Add a `streaming:` block to your serve config.")
sessions = SessionManager(
segment_cap=serve_config.streaming.generation_segment_cap,
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
)
state = ServerState(
serve_config=serve_config,
generator=generator,
sessions=sessions,
session_store=session_store or InMemorySessionStore(),
)
app = FastAPI(title="FastVideo Streaming")
@app.get("/health")
async def _health() -> JSONResponse:
return JSONResponse({
"status": "ok",
"sessions": len(state.sessions),
"stream_mode": state.serve_config.streaming.stream_mode,
})
@app.websocket("/v1/stream")
async def _stream(websocket: WebSocket) -> None:
await websocket.accept()
try:
session = state.sessions.create()
except SessionRejected as exc:
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
return
try:
await _handle_session(websocket, session, state)
except WebSocketDisconnect:
logger.info("session %s: client disconnected", session.id[:8])
except Exception: # pragma: no cover - defensive catch-all
logger.exception("session %s: unhandled error", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
finally:
_cleanup_session(session, state)
app.state.server_state = state
return app
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
"""Launch the streaming server.
Boots a :class:`fastvideo.VideoGenerator` from
``serve_config.generator`` unless ``generator`` is provided, then
serves ``build_app(...)`` via uvicorn.
"""
def run_server(serve_config: ServeConfig) -> None:
"""Launch the streaming (WebSocket / Dynamo) server."""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
"got None. Add a `streaming:` block to your serve config.")
import uvicorn
if generator is None:
from fastvideo import VideoGenerator # lazy to avoid boot cost
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
app = build_app(serve_config, generator)
uvicorn.run(
app,
host=serve_config.server.host,
port=serve_config.server.port,
)
async def _handle_session(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
init = await _read_init_message(websocket, session, state)
if init is None:
return
await _apply_session_init(session, init, state)
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
session.transition(SessionState.GPU_BINDING)
await _send_json(websocket, GpuAssigned(
gpu_id=0,
session_timeout=state.sessions.session_timeout_seconds,
))
session.transition(SessionState.ACTIVE)
await _send_json(websocket, _build_stream_start(session, state))
try:
await _run_segment_loop(websocket, session, state)
finally:
with contextlib.suppress(RuntimeError):
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
async def _read_init_message(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> SessionInitV2 | None:
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: init timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return None
except WebSocketDisconnect:
return None
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
return None
if not isinstance(parsed, SessionInitV2):
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
return None
return parsed
async def _reject_init(
websocket: WebSocket,
session: Session,
message: str,
close_reason: str,
) -> None:
await _send_error(websocket, "invalid_message", message, retryable=False)
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.REJECTED)
async def _apply_session_init(
session: Session,
init: SessionInitV2,
state: ServerState,
) -> None:
session.client_id = init.client_id
session.preset = init.preset
session.preset_label = init.preset_label
session.curated_prompts = list(init.curated_prompts)
session.enhancement_enabled = init.enhancement_enabled
session.auto_extension_enabled = init.auto_extension_enabled
session.loop_generation_enabled = init.loop_generation_enabled
session.single_clip_mode = init.single_clip_mode
session.stream_mode = init.stream_mode
if init.initial_image is not None:
# Decode + disk write off the event loop; payload is up to 32 MiB.
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
if image is not None:
session.metadata["session_init_image"] = image.path
if init.continuation_state is not None:
session.continuation_state = _coerce_state(init.continuation_state)
if session.continuation_state is not None:
state.session_store.store(session.id, session.continuation_state)
async def _run_segment_loop(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
cap = state.sessions.segment_cap
while True:
if session.segment_cap_reached(cap):
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
return
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: idle timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return
except WebSocketDisconnect:
return
session.touch()
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
continue
if isinstance(parsed, SnapshotState):
snap = state.session_store.snapshot(session.id)
if snap is None:
await _send_error(websocket,
"internal_error",
"no continuation state available for session",
retryable=False)
continue
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
continue
if isinstance(parsed, SegmentPromptSource):
await _run_segment(websocket, session, state, parsed)
continue
# Silently ignore unknown-but-valid types (additive-evolution
# rule in streaming.md).
_apply_toggle(session, parsed)
async def _run_segment(
websocket: WebSocket,
session: Session,
state: ServerState,
message: SegmentPromptSource,
) -> None:
request = _build_generation_request(session, message, state)
segment_idx = session.segment_idx
await _send_json(
websocket,
Ltx2SegmentStart(
segment_idx=segment_idx,
prompt=message.prompt,
total_steps=request.sampling.num_inference_steps,
))
start = time.perf_counter()
loop = asyncio.get_running_loop()
# TODO: executor-wrapped generate() cannot be cancelled, so a
# client disconnect mid-segment leaves the GPU work running to
# completion. Real cancellation needs the generate_async API.
try:
result = await loop.run_in_executor(None, state.generator.generate, request)
except Exception as exc:
logger.exception("session %s: generator failed", session.id[:8])
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
elapsed_ms = (time.perf_counter() - start) * 1000.0
frames = _extract_frames(result)
if not frames:
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
# Synchronous generator call has no per-step hook; emit one
# terminal StepComplete so observability wiring still sees the
# segment finish.
total = request.sampling.num_inference_steps
await _send_json(websocket, StepComplete(
segment_idx=segment_idx,
step=total,
total_steps=total,
stage="denoise",
))
encoder = FragmentedMP4Encoder(
width=request.sampling.width,
height=request.sampling.height,
fps=request.sampling.fps,
segment_idx=segment_idx,
)
chunks_relayed = 0
async with encoder:
init_sent = False
async for chunk in encoder.encode(frames):
if chunk.kind == "init":
await _send_json(websocket, MediaInit(
segment_idx=segment_idx,
stream_id=chunk.stream_id,
))
init_sent = True
await websocket.send_bytes(chunk.data)
if init_sent and chunk.kind == "media":
chunks_relayed += 1
await _send_json(
websocket,
MediaSegmentComplete(
segment_idx=segment_idx,
stream_id=encoder.stream_id,
chunks=chunks_relayed,
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
))
new_state = _extract_state(result)
if new_state is not None:
session.continuation_state = new_state
state.session_store.store(session.id, new_state)
session.segment_idx += 1
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ACTIVE)
await _send_json(
websocket,
Ltx2SegmentComplete(
segment_idx=segment_idx,
generation_time_ms=elapsed_ms,
e2e_latency_ms=elapsed_ms,
))
def _build_stream_start(
session: Session,
state: ServerState,
) -> Ltx2StreamStart:
default = state.serve_config.default_request
return Ltx2StreamStart(
preset=session.preset,
width=default.sampling.width,
height=default.sampling.height,
fps=default.sampling.fps,
num_frames=default.sampling.num_frames,
)
def _build_generation_request(
session: Session,
message: SegmentPromptSource,
state: ServerState,
) -> GenerationRequest:
# Start from the operator-pinned default_request to pick up the
# preset-selected sampling knobs; override with per-message values.
base = state.serve_config.default_request
sampling_kwargs: dict[str, Any] = {
"num_videos_per_prompt":
base.sampling.num_videos_per_prompt,
"seed":
message.seed if message.seed is not None else base.sampling.seed,
"num_frames":
base.sampling.num_frames,
"height":
base.sampling.height,
"width":
base.sampling.width,
"fps":
base.sampling.fps,
"num_inference_steps":
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
"guidance_scale":
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
}
request = GenerationRequest(
prompt=message.prompt,
negative_prompt=message.negative_prompt or base.negative_prompt,
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
sampling=SamplingConfig(**sampling_kwargs),
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
state=session.continuation_state,
)
return request
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
kind = raw.get("kind")
payload = raw.get("payload")
if not isinstance(kind, str) or not isinstance(payload, dict):
return None
return ContinuationState(kind=kind, payload=payload)
def _apply_toggle(session: Session, message: Any) -> None:
if isinstance(message, EnhancementUpdated):
session.enhancement_enabled = message.enabled
elif isinstance(message, AutoExtensionUpdated):
session.auto_extension_enabled = message.enabled
elif isinstance(message, LoopGenerationUpdated):
session.loop_generation_enabled = message.enabled
elif isinstance(message, GenerationPausedUpdated):
session.generation_paused = message.paused
elif isinstance(message, SeedPromptsUpdated):
session.curated_prompts = list(message.seed_prompts)
def _extract_frames(result: Any) -> list:
if hasattr(result, "frames"):
return list(result.frames or [])
if isinstance(result, dict):
return list(result.get("frames") or [])
return []
def _extract_state(result: Any) -> ContinuationState | None:
state = getattr(result, "state", None)
if state is None and isinstance(result, dict):
state = result.get("state")
if isinstance(state, ContinuationState):
return state
if isinstance(state, dict):
return _coerce_state(state)
return None
async def _send_json(websocket: WebSocket, message: Any) -> None:
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
await websocket.send_json(payload)
async def _send_error(
websocket: WebSocket,
code: str,
message: str,
*,
retryable: bool,
) -> None:
await _send_json(
websocket,
ErrorMessage(code=code, message=message, retryable=retryable),
)
def _cleanup_session(session: Session, state: ServerState) -> None:
state.sessions.close(session.id)
state.session_store.drop(session.id)
init_image_path = session.metadata.get("session_init_image")
if isinstance(init_image_path, str):
with contextlib.suppress(FileNotFoundError):
os.unlink(init_image_path)
__all__ = [
"ServerState",
"build_app",
"run_server",
]
raise NotImplementedError("streaming server is not implemented yet")
-214
View File
@@ -1,214 +0,0 @@
# 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",
]
@@ -1,103 +0,0 @@
# 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",
]
@@ -1,206 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Session state store for the FastVideo streaming server.
The streaming server keeps continuation state (decoded frames + audio
latents from the previous segment) server-side so the client doesn't
re-upload multi-megabyte tensors each WebSocket message. Two operations
are needed:
* ``snapshot(session_id) -> ContinuationState`` — serialize the current
state so it can be exported (e.g. over HTTP) or migrated to a
different server.
* ``hydrate(state) -> session_id`` — load a previously serialized state
into a new session (for resume-after-disconnect flows).
The store is an ABC with an :class:`InMemorySessionStore` default; Redis
or other backends can drop in without touching the pipeline.
Large tensor payloads (video frames, audio latents) are kept out of the
JSON payload via an accompanying :class:`BlobStore`. Both stores share a
process today; they are separate types so that a future implementation
can put blobs on S3 while keeping session metadata in Redis.
"""
from __future__ import annotations
import threading
import uuid
from abc import ABC, abstractmethod
from collections.abc import Iterator
from dataclasses import dataclass
from fastvideo.api.schema import ContinuationState
class BlobStore(ABC):
"""Opaque byte-blob storage keyed by id.
A :class:`ContinuationState` payload can reference large tensors
stored in a :class:`BlobStore` rather than inlining them, so the
JSON payload stays small when the state travels over the wire.
"""
@abstractmethod
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
"""Store ``data`` and return a blob id for later retrieval."""
@abstractmethod
def get(self, blob_id: str) -> bytes:
"""Load a previously stored blob. Raises ``KeyError`` if absent."""
@abstractmethod
def drop(self, blob_id: str) -> None:
"""Remove a blob. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, blob_id: str) -> bool:
...
@dataclass(frozen=True)
class _BlobRecord:
data: bytes
mime: str
class InMemoryBlobStore(BlobStore):
"""Thread-safe in-memory :class:`BlobStore` for single-process servers.
No eviction policy — callers are responsible for calling
:meth:`drop` when a blob's owning state is replaced or a session
ends. A redis- or filesystem-backed :class:`BlobStore` should
replace this when the streaming server lands as a real service
(PR 7.5+).
"""
def __init__(self) -> None:
self._blobs: dict[str, _BlobRecord] = {}
self._lock = threading.Lock()
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
blob_id = uuid.uuid4().hex
with self._lock:
self._blobs[blob_id] = _BlobRecord(data=data, mime=mime)
return blob_id
def get(self, blob_id: str) -> bytes:
with self._lock:
record = self._blobs.get(blob_id)
if record is None:
raise KeyError(f"Unknown blob id: {blob_id}")
return record.data
def drop(self, blob_id: str) -> None:
with self._lock:
self._blobs.pop(blob_id, None)
def __contains__(self, blob_id: str) -> bool:
with self._lock:
return blob_id in self._blobs
def __len__(self) -> int:
with self._lock:
return len(self._blobs)
class SessionStore(ABC):
"""Keyed store for per-session continuation state.
Implementations own the session-id → state mapping. The streaming
server calls :meth:`store` after each segment and :meth:`snapshot`
when a client explicitly asks for an exportable state handle.
"""
@abstractmethod
def store(self, session_id: str, state: ContinuationState) -> None:
"""Persist ``state`` for ``session_id``, replacing any prior value."""
@abstractmethod
def snapshot(self, session_id: str) -> ContinuationState | None:
"""Return the current state for ``session_id`` (or ``None``)."""
@abstractmethod
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
"""Install ``state`` as the starting point for a session.
When ``session_id`` is ``None`` the store allocates a fresh id
(UUID4); when provided the store uses it verbatim, overwriting
any prior state at that id.
"""
@abstractmethod
def drop(self, session_id: str) -> None:
"""Forget a session. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, session_id: str) -> bool:
...
@abstractmethod
def __iter__(self) -> Iterator[str]:
...
class InMemorySessionStore(SessionStore):
"""Thread-safe in-memory :class:`SessionStore`.
Default implementation used by single-process deployments; a future
Redis-backed store can be dropped in without changes to the server.
No eviction / TTL / bounded capacity — sessions only leave via
:meth:`drop`. The live streaming server (PR 7.5+) is responsible
for bounding growth and for dropping any :class:`BlobStore` blobs
referenced by a state when that state is replaced or a session
ends; this class does not know about blobs.
"""
def __init__(self) -> None:
self._sessions: dict[str, ContinuationState] = {}
self._lock = threading.Lock()
def store(self, session_id: str, state: ContinuationState) -> None:
with self._lock:
self._sessions[session_id] = state
def snapshot(self, session_id: str) -> ContinuationState | None:
with self._lock:
return self._sessions.get(session_id)
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
sid = session_id or uuid.uuid4().hex
with self._lock:
self._sessions[sid] = state
return sid
def drop(self, session_id: str) -> None:
with self._lock:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
with self._lock:
return session_id in self._sessions
def __iter__(self) -> Iterator[str]:
with self._lock:
return iter(list(self._sessions))
def __len__(self) -> int:
with self._lock:
return len(self._sessions)
__all__ = [
"BlobStore",
"InMemoryBlobStore",
"InMemorySessionStore",
"SessionStore",
]
-213
View File
@@ -1,213 +0,0 @@
# 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",
]
+1 -11
View File
@@ -844,13 +844,6 @@ 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
@@ -1111,10 +1104,7 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--real-score-guidance-scale",
type=float,
default=TrainingArgs.real_score_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)."))
help="Teacher guidance scale")
parser.add_argument("--fake-score-learning-rate",
type=float,
default=TrainingArgs.fake_score_learning_rate,
@@ -1,5 +0,0 @@
# 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
@@ -1,386 +0,0 @@
# 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",
]
+2 -2
View File
@@ -2,7 +2,7 @@
"""LTX2 model family pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_stage_override_fields, )
REFINE_STAGE_OVERRIDE_FIELDS, )
_LTX2_NEGATIVE_PROMPT = ("blurry, out of focus, overexposed, underexposed, low contrast, "
"washed out colors, excessive noise, grainy texture, poor lighting, "
@@ -36,7 +36,7 @@ _REFINE_STAGE = PresetStageSpec(
name="refine",
kind="refinement",
description="Latent-upsample + second-pass refine",
allowed_overrides=refine_stage_override_fields(),
allowed_overrides=REFINE_STAGE_OVERRIDE_FIELDS,
)
LTX2_BASE = InferencePreset(
@@ -27,8 +27,6 @@ class LTX2RefinePresetOverride:
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
@@ -42,18 +40,15 @@ def refine_override_to_dict(override: LTX2RefinePresetOverride | LTX2RefineStage
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))
REFINE_PRESET_OVERRIDE_FIELDS: frozenset[str] = frozenset(f.name for f in fields(LTX2RefinePresetOverride))
REFINE_STAGE_OVERRIDE_FIELDS: frozenset[str] = frozenset(f.name for f in fields(LTX2RefineStageOverride))
REFINE_FLAT_KEYS: frozenset[str] = (REFINE_PRESET_OVERRIDE_FIELDS | REFINE_STAGE_OVERRIDE_FIELDS)
__all__ = [
"LTX2RefinePresetOverride",
"LTX2RefineStageOverride",
"REFINE_FLAT_KEYS",
"REFINE_PRESET_OVERRIDE_FIELDS",
"REFINE_STAGE_OVERRIDE_FIELDS",
"refine_override_to_dict",
"refine_preset_override_fields",
"refine_stage_override_fields",
]
@@ -26,7 +26,7 @@ MATRIXGAME_I2V = InferencePreset(
"fps": 25,
"guidance_scale": 1.0,
"num_inference_steps": 3,
"negative_prompt": "",
"negative_prompt": None,
},
)
@@ -26,7 +26,7 @@ TURBO_T2V_1_3B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": "",
"negative_prompt": None,
},
)
@@ -44,7 +44,7 @@ TURBO_T2V_14B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": "",
"negative_prompt": None,
},
)
@@ -62,7 +62,7 @@ TURBO_I2V_A14B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": "",
"negative_prompt": None,
},
)
@@ -17,8 +17,6 @@ import torch
if TYPE_CHECKING:
from torchcodec.decoders import VideoDecoder
from fastvideo.api.schema import ContinuationState
import time
from collections import OrderedDict
@@ -208,9 +206,6 @@ 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)
@@ -524,22 +524,6 @@ 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:
+4 -16
View File
@@ -519,7 +519,7 @@ def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch)
cli_main.main()
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, monkeypatch):
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
config_path = tmp_path / "serve-streaming.yaml"
config_path.write_text(
"generator:\n"
@@ -530,21 +530,9 @@ def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, mo
)
args, _ = _parse_serve_args(["--config", str(config_path)])
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"
with pytest.raises(NotImplementedError,
match="streaming server is not implemented"):
ServeSubcommand().cmd(args)
def test_streaming_run_server_rejects_missing_streaming_block():
@@ -155,51 +155,6 @@ class TestLegacyLtx2VaeTilingTranslation:
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
# -------------------------------------------------------------------
@@ -1,250 +0,0 @@
# 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
@@ -4,14 +4,8 @@
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.
public typed ``GeneratorConfig`` surface can represent it end-to-end,
with no fields silently falling through to ``pipeline.experimental``.
"""
from __future__ import annotations
@@ -27,19 +21,10 @@ from fastvideo.api.compat import (
# 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.
# Skipped here because they are opaque Python objects that legitimately
# belong in experimental:
# - pipeline_config=<PipelineConfig instance>
# - enable_torch_compile_text_encoder (not in public FastVideoArgs)
GPU_POOL_LOAD_KWARGS = {
"config_model_path": "/models/ltx2-distilled/config",
"num_gpus": 1,
@@ -57,7 +42,6 @@ GPU_POOL_LOAD_KWARGS = {
"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,
@@ -94,7 +78,6 @@ class TestGpuPoolForwardTranslation:
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"
@@ -133,11 +116,9 @@ class TestGpuPoolForwardTranslation:
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.
"""
original gpu_pool flat-kwarg shape, so callers can wire a public
``gpu_pool`` through ``generator_config_to_fastvideo_args`` without
the runtime noticing."""
@pytest.fixture
def args_kwargs(self, monkeypatch):
@@ -187,12 +168,6 @@ class TestGpuPoolReverseTranslation:
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
@@ -213,9 +188,7 @@ class TestRefineFlattenCoversAllTypedFields:
)
from fastvideo.api.schema import GeneratorConfig, PipelineSelection
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_preset_override_fields,
refine_stage_override_fields,
)
REFINE_FLAT_KEYS, )
captured: dict[str, object] = {}
@@ -240,9 +213,7 @@ class TestRefineFlattenCoversAllTypedFields:
"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, (
assert set(refine_payload) == REFINE_FLAT_KEYS, (
"payload must cover every typed field to exercise the flatten loop")
config = GeneratorConfig(
@@ -9,9 +9,9 @@ from fastvideo.api.presets import get_preset, validate_stage_overrides
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
LTX2RefinePresetOverride,
LTX2RefineStageOverride,
REFINE_PRESET_OVERRIDE_FIELDS,
REFINE_STAGE_OVERRIDE_FIELDS,
refine_override_to_dict,
refine_preset_override_fields,
refine_stage_override_fields,
)
@@ -57,7 +57,7 @@ class TestRefineStageOverrideDataclass:
}
def test_fields_accessor_matches_dataclass(self) -> None:
assert refine_stage_override_fields() == frozenset({
assert REFINE_STAGE_OVERRIDE_FIELDS == frozenset({
"num_inference_steps",
"guidance_scale",
"image_crf",
@@ -86,7 +86,7 @@ class TestRefinePresetOverrideDataclass:
}
def test_fields_accessor_matches_dataclass(self) -> None:
assert refine_preset_override_fields() == frozenset({
assert REFINE_PRESET_OVERRIDE_FIELDS == frozenset({
"enabled",
"add_noise",
})
@@ -101,7 +101,7 @@ class TestStageOverridesMirrorPresetSchema:
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()
assert refine_schema.allowed_overrides == REFINE_STAGE_OVERRIDE_FIELDS
def test_roundtrip_through_validate_stage_overrides(self) -> None:
import fastvideo.registry # noqa: F401
-1
View File
@@ -113,7 +113,6 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
},
"compile": {
"enabled": False,
"text_encoder_enabled": None,
"backend": None,
"fullgraph": None,
"mode": None,
-37
View File
@@ -585,40 +585,3 @@ 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))
@@ -71,7 +71,13 @@ def _get_extra_dataclass_fields(
continue
for _, modname, is_pkg in pkgutil.walk_packages(
package.__path__, prefix=f"{package_name}."):
if modname.endswith(".__pycache__"):
# Flat ``configs.pipelines.<family>`` modules carry the config
# directly; colocated ``basic.<family>.pipeline_configs``
# submodules do too. Everything else under ``basic`` is heavy
# model code we don't need to import for a schema check.
basename = modname.rsplit(".", 1)[-1]
is_flat = modname.startswith("fastvideo.configs.pipelines.")
if not is_flat and basename != "pipeline_configs":
continue
module = importlib.import_module(modname)
for obj in vars(module).values():
@@ -1 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
@@ -1,154 +0,0 @@
# 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"
@@ -1,237 +0,0 @@
# 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"])
@@ -1,133 +0,0 @@
# 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()
@@ -1,90 +0,0 @@
# 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
@@ -1,186 +0,0 @@
# 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)
@@ -1,99 +0,0 @@
# 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())
+110 -110
View File
@@ -10,21 +10,21 @@ Note: num_inference_steps is reduced to 4 for faster CI.
import os
import pytest
import torch
from fastvideo import VideoGenerator
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.logger import init_logger
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
from fastvideo.tests.ssim.reference_utils import (
build_generated_output_dir,
build_reference_folder_path,
get_cuda_device_name,
resolve_device_reference_folder,
select_ssim_params,
)
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
import pytest
import torch
from fastvideo import VideoGenerator
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.logger import init_logger
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
from fastvideo.tests.ssim.reference_utils import (
build_generated_output_dir,
build_reference_folder_path,
get_cuda_device_name,
resolve_device_reference_folder,
select_ssim_params,
)
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
logger = init_logger(__name__)
@@ -48,20 +48,20 @@ def _find_lingbotworld_examples_root() -> str | None:
return None
device_name = get_cuda_device_name()
device_reference_folder = resolve_device_reference_folder(
(
("A40", "A40"),
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
),
device_name=device_name,
logger=logger,
)
device_name = get_cuda_device_name()
device_reference_folder = resolve_device_reference_folder(
(
("A40", "A40"),
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
),
device_name=device_name,
logger=logger,
)
LINGBOT_PARAMS = {
LINGBOT_PARAMS = {
"model_path": "FastVideo/LingBot-World-Base-Cam-Diffusers",
"num_gpus": 2,
"height": 256,
@@ -88,29 +88,29 @@ LINGBOT_PARAMS = {
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
"皮肤,肢体,面部特征,汽车,电线"
),
}
_LINGBOT_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
LINGBOT_PARAMS["model_path"])
LINGBOT_FULL_QUALITY_PARAMS = {
"model_path": LINGBOT_PARAMS["model_path"],
"num_gpus": LINGBOT_PARAMS["num_gpus"],
"height": _LINGBOT_FULL_QUALITY_DEFAULTS.height,
"width": _LINGBOT_FULL_QUALITY_DEFAULTS.width,
"num_frames": LINGBOT_PARAMS["num_frames"], # default num_frames: 125
"num_inference_steps": _LINGBOT_FULL_QUALITY_DEFAULTS.num_inference_steps,
"guidance_scale": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale,
"guidance_scale_2": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale_2,
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
"flow_shift": LINGBOT_PARAMS["flow_shift"],
"boundary_ratio": _LINGBOT_FULL_QUALITY_DEFAULTS.boundary_ratio,
"seed": _LINGBOT_FULL_QUALITY_DEFAULTS.seed,
"fps": _LINGBOT_FULL_QUALITY_DEFAULTS.fps,
"spatial_scale": LINGBOT_PARAMS["spatial_scale"],
"example_case": LINGBOT_PARAMS["example_case"],
"image_path": LINGBOT_PARAMS["image_path"],
"negative_prompt": _LINGBOT_FULL_QUALITY_DEFAULTS.negative_prompt,
}
),
}
_LINGBOT_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
LINGBOT_PARAMS["model_path"])
LINGBOT_FULL_QUALITY_PARAMS = {
"model_path": LINGBOT_PARAMS["model_path"],
"num_gpus": LINGBOT_PARAMS["num_gpus"],
"height": _LINGBOT_FULL_QUALITY_DEFAULTS.height,
"width": _LINGBOT_FULL_QUALITY_DEFAULTS.width,
"num_frames": LINGBOT_PARAMS["num_frames"], # default num_frames: 125
"num_inference_steps": _LINGBOT_FULL_QUALITY_DEFAULTS.num_inference_steps,
"guidance_scale": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale,
"guidance_scale_2": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale_2,
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
"flow_shift": LINGBOT_PARAMS["flow_shift"],
"boundary_ratio": _LINGBOT_FULL_QUALITY_DEFAULTS.boundary_ratio,
"seed": _LINGBOT_FULL_QUALITY_DEFAULTS.seed,
"fps": _LINGBOT_FULL_QUALITY_DEFAULTS.fps,
"spatial_scale": LINGBOT_PARAMS["spatial_scale"],
"example_case": LINGBOT_PARAMS["example_case"],
"image_path": LINGBOT_PARAMS["image_path"],
"negative_prompt": _LINGBOT_FULL_QUALITY_DEFAULTS.negative_prompt,
}
TEST_PROMPTS = [
"The video presents a soaring journey through a fantasy jungle. The wind "
@@ -123,80 +123,80 @@ TEST_PROMPTS = [
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
params = select_ssim_params(LINGBOT_PARAMS, LINGBOT_FULL_QUALITY_PARAMS)
if device_reference_folder is None:
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
if torch.cuda.device_count() < params["num_gpus"]:
pytest.skip(
f"LingBot SSIM test requires {params['num_gpus']} GPUs, "
f"but only {torch.cuda.device_count()} detected."
)
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
params = select_ssim_params(LINGBOT_PARAMS, LINGBOT_FULL_QUALITY_PARAMS)
if device_reference_folder is None:
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
if torch.cuda.device_count() < params["num_gpus"]:
pytest.skip(
f"LingBot SSIM test requires {params['num_gpus']} GPUs, "
f"but only {torch.cuda.device_count()} detected."
)
examples_root = _find_lingbotworld_examples_root()
if examples_root is None:
pytest.skip(
"lingbotworld_examples not found under examples/inference/basic.")
action_path = os.path.join(examples_root, params["example_case"])
action_path = os.path.join(examples_root, params["example_case"])
if not (os.path.exists(os.path.join(action_path, "poses.npy"))
and os.path.exists(os.path.join(action_path, "intrinsics.npy"))):
pytest.skip(f"Missing camera npy files under {action_path}")
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=params["num_frames"],
height=params["height"],
width=params["width"],
spatial_scale=params["spatial_scale"],
)
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=params["num_frames"],
height=params["height"],
width=params["width"],
spatial_scale=params["spatial_scale"],
)
script_dir = os.path.dirname(os.path.abspath(__file__))
model_id = "LingBot-World-Base-Cam-Diffusers"
output_dir = build_generated_output_dir(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
output_dir = build_generated_output_dir(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
init_kwargs = {
"num_gpus": params["num_gpus"],
"flow_shift": params["flow_shift"],
"boundary_ratio": params["boundary_ratio"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
init_kwargs = {
"num_gpus": params["num_gpus"],
"flow_shift": params["flow_shift"],
"boundary_ratio": params["boundary_ratio"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": True,
"vae_cpu_offload": False,
"pin_cpu_memory": True,
}
generation_kwargs = {
"output_path": output_dir,
"image_path": params["image_path"],
"height": params["height"],
"width": params["width"],
"num_frames": aligned_num_frames,
"num_inference_steps": params["num_inference_steps"],
"guidance_scale": params["guidance_scale"],
"guidance_scale_2": params["guidance_scale_2"],
"embedded_cfg_scale": params["embedded_cfg_scale"],
"seed": params["seed"],
"fps": params["fps"],
"negative_prompt": params["negative_prompt"],
"c2ws_plucker_emb": c2ws_plucker_emb,
}
generation_kwargs = {
"output_path": output_dir,
"image_path": params["image_path"],
"height": params["height"],
"width": params["width"],
"num_frames": aligned_num_frames,
"num_inference_steps": params["num_inference_steps"],
"guidance_scale": params["guidance_scale"],
"guidance_scale_2": params["guidance_scale_2"],
"embedded_cfg_scale": params["embedded_cfg_scale"],
"seed": params["seed"],
"fps": params["fps"],
"negative_prompt": params["negative_prompt"],
"c2ws_plucker_emb": c2ws_plucker_emb,
}
generator: VideoGenerator | None = None
try:
generator = VideoGenerator.from_pretrained(
model_path=params["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
generator = VideoGenerator.from_pretrained(
model_path=params["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
finally:
if generator is not None:
generator.shutdown()
@@ -205,12 +205,12 @@ def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
assert os.path.exists(generated_video_path), (
f"Output video was not generated at {generated_video_path}")
reference_folder = build_reference_folder_path(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
reference_folder = build_reference_folder_path(
script_dir,
device_reference_folder,
model_id,
ATTENTION_BACKEND,
)
if not os.path.exists(reference_folder):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
@@ -234,11 +234,11 @@ def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
mean_ssim = ssim_values[0]
logger.info("SSIM mean value: %s", mean_ssim)
write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path,
params["num_inference_steps"], prompt)
write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path,
params["num_inference_steps"], prompt)
min_acceptable_ssim = 0.70
min_acceptable_ssim = 0.90
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}")
+1 -1
View File
@@ -89,5 +89,5 @@ def test_ltx2_distilled_inference_similarity(
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,
min_acceptable_ssim=0.98,
)
@@ -145,7 +145,6 @@ 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",
@@ -1,91 +0,0 @@
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,10 +666,6 @@ 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,52 +294,3 @@ 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