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
156 changed files with 853 additions and 12515 deletions
@@ -29,8 +29,8 @@
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
],
"run_config": {
"num_warmup_runs": 2,
"num_measurement_runs": 5,
"num_warmup_runs": 1,
"num_measurement_runs": 3,
"required_gpus": 2
},
"thresholds": {
+2 -73
View File
@@ -63,72 +63,7 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
EFFECTIVE_PR=$PR_NUMBER
fi
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} IMAGE_VERSION=$IMAGE_VERSION"
POST_RUN_HOOK=""
upload_performance_artifacts() {
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
LOCAL_DIR="downloaded_reports"
_download_reports() {
log "Downloading perf_reports/ from Modal Volume..."
mkdir -p "$LOCAL_DIR"
if ! modal volume get hf-model-weights "perf_reports/" "$LOCAL_DIR"; then
log "Error: Failed to download perf_reports/ from Modal Volume."
return 1
fi
}
_upload_dashboard() {
local target
target=$(find "$LOCAL_DIR" -name "dashboard_*${SHORT_SHA}*" | head -n 1)
log "TARGET dashboard: '$target'"
if [ -n "$target" ]; then
log "Found dashboard: $target. Uploading to Buildkite..."
buildkite-agent artifact upload "$target"
buildkite-agent annotate --style info --context "perf-dashboard" < "$target"
else
log "Warning: Could not find a dashboard file matching $SHORT_SHA"
fi
}
_upload_perf_summary() {
local target
target=$(find "$LOCAL_DIR" -name "perf_*${SHORT_SHA}*" | head -n 1)
log "TARGET perf summary: '$target'"
if [ -n "$target" ]; then
log "Found perf summary: $target. Uploading to Buildkite..."
buildkite-agent artifact upload "$target"
buildkite-agent annotate --style info --context "perf-summary" < "$target"
else
log "Warning: Could not find a perf summary file matching $SHORT_SHA"
fi
}
_cleanup_modal_volume() {
log "Cleaning up perf_reports/ from Modal Volume..."
if modal volume rm hf-model-weights "perf_reports/" --recursive; then
log "Successfully deleted perf_reports/ from Modal Volume."
else
log "Warning: Failed to delete perf_reports/ from Modal Volume. Manual cleanup may be required."
fi
}
_cleanup_local() {
log "Cleaning up local download directory..."
rm -rf "$LOCAL_DIR"
}
# --- Main flow ---
_download_reports || { _cleanup_local; return 1; }
_upload_dashboard
_upload_perf_summary
_cleanup_modal_volume
_cleanup_local
}
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
@@ -189,9 +124,8 @@ case "$TEST_TYPE" in
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
;;
"performance")
log "Running performance tests on Modal..."
log "Running performance tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_performance_tests"
POST_RUN_HOOK="upload_performance_artifacts"
;;
"api_server")
log "Running API server integration tests..."
@@ -213,10 +147,5 @@ else
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
fi
if [ -n "$POST_RUN_HOOK" ]; then
log "Executing post-run hook: $POST_RUN_HOOK"
"$POST_RUN_HOOK"
fi
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
exit $TEST_EXIT_CODE
-1
View File
@@ -18,7 +18,6 @@ exclude: |
fastvideo/train\.py|
fastvideo/utils/.*|
examples/.*|
\.agents/.*|
.github/workflows/publish-fastvideo.yml|
.github/workflows/_template-build-image.yml|
docs/source/inference/support_matrix.md
+3 -6
View File
@@ -296,10 +296,8 @@ Action:
- Add or reuse a numerical parity test that loads the official model and the
FastVideo model and compares outputs.
- See examples in `tests/local_tests/` organized by model family
(e.g., `tests/local_tests/sd35/`, `tests/local_tests/ltx2/`,
`tests/local_tests/stable_audio/`) and the navigation index in
`tests/local_tests/README.md`.
- See examples in `tests/local_tests/` (e.g., `tests/local_tests/upsamplers/`)
and the commands in `tests/local_tests/README.md`.
- If there are discrepancies, add opt‑in logging to both models and compare
activation summaries (layer output sums, per‑stage logs).
- First align the loaded weights (validate `param_names_mapping`).
@@ -350,8 +348,7 @@ Purpose:
Action:
- Add a pipeline parity test under `tests/local_tests/<family>/`
(e.g., `tests/local_tests/<family>/test_<family>_pipeline_parity.py`).
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
- See the [Testing Guide](testing.md) for test conventions.
### 7) Add user‑facing examples
@@ -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.
@@ -1,77 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 — text-to-audio (baseline) example.
User story (game-audio designer, prototyping):
"I'm prototyping a level and I need 6 seconds of background
ambience — gentle wind, distant thunder, a hint of birdsong. I
don't want to dig through a sound library; I want to type what I
hear in my head and get a wav back. If it's wrong I'll iterate
on the prompt. This is the first stop."
User story (musician sketching ideas):
"I want to bounce a 30s lo-fi drum loop to use as a placeholder
bed while I build the rest of the track. Type prompt, get audio,
drop into the DAW. The actual production beat I'll record
myself, but I need *something* to write the chords against."
User story (researcher exploring the model):
"First time touching Stable Audio Open — what does it sound
like at default settings? This is the smallest amount of code
that goes from prompt to mp4."
How it works:
Pure text-to-audio (T2A). The pipeline runs:
T5 + NumberConditioner -> StableAudioDiT -> Oobleck VAE
via the `dpmpp-3m-sde` k-diffusion sampler. All components are
FastVideo-native — no diffusers / transformers model imports at
runtime (see REVIEW item 30). Mirrors upstream
`stable_audio_tools.inference.generation.generate_diffusion_cond`
bit-for-bit (~0.2% abs_mean drift on 25 steps).
Tunable knobs (the "creative dials"):
audio_end_in_s
1–6 — quick ideation (sub-10s wall clock at 100 steps)
10–30 — full musical phrase / loop length (the README example
uses 30s)
47.5 — model maximum (full sample_size = 2097152 / 44100 Hz)
num_inference_steps
25 — fast preview, occasional artifacts
100 — preset default (matches the HF model card)
250 — diminishing returns past here
guidance_scale
3 — looser, more variation per seed
7 — preset default; matches README
12+ — sharper but can sound "fried"
Prerequisites:
1. Accept the terms on https://huggingface.co/stabilityai/stable-audio-open-1.0
and export your HF token in the shell:
export HF_TOKEN=hf_...
2. Install optional inference deps (one-time):
pip install k_diffusion einops_exts alias_free_torch torchsde
"""
from fastvideo import VideoGenerator
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/stable-audio-open-1.0-Diffusers",
num_gpus=1,
)
output_path = "outputs_audio/stable_audio_basic/output_stable_audio.wav"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
# 6-second clip; the model max is ~47.5s.
audio_end_in_s=6.0,
# The registered preset gives 100 steps + CFG=7.0 by default;
# override num_inference_steps / guidance_scale here for QA.
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,77 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 — audio-to-audio variation example.
User story (musician, late at night):
"I generated this 12-second lo-fi loop earlier and I love the chord
progression and overall vibe, but the snare hit at 0:08 sounds wrong
and the rhythm feels stiff. I don't want to start over from scratch
and lose what's working — I want the model to keep the harmony and
mood but reroll the percussion + groove."
User story (sound designer, on a deadline):
"I have one good 'sword clang' SFX. The art director wants 8 sibling
variations that all feel like the same sword from different angles —
same metal, same weight, slightly different impact. I'd rather
refine my one good take than text-prompt my way through 50 misses."
Pass `init_audio=path/to/clip` (any wav/mp3/mp4/m4a/flac the standard
deps decode) and the model will use it as a starting point for the
text prompt instead of pure noise.
Picking `init_audio_strength` (0.0 to 1.0):
Higher = closer to the source clip. Lower = more transformation.
(Same convention as the "Input Audio Strength" slider in
Stability's commercial Stable Audio web UI, so values transfer
directly.)
| strength | what you get |
|----------|----------------------------------------------------|
| 1.00 | Output ≈ reference. No transformation. |
| 0.85 | Texture micro-variation only. |
| 0.70 | Light reroll, same instruments. |
| 0.60 | Default. Instrument identity is replaceable |
| | (cello can take over from piano on the same notes).|
| 0.50 | Heavy — only melody / chord progression survives. |
| 0.30 | Reference acts as a loose mood prompt. |
| 0.00 | Plain T2A — reference ignored. |
Rule of thumb by intent:
* "Fix one part of this clip" -> 0.75 .. 0.85
* "Same notes, different instrument" -> 0.55 .. 0.65
* "Same chord progression, new content" -> 0.40 .. 0.55
* "Use this as a loose mood prompt" -> 0.20 .. 0.35
If the reference timbre is bleeding through more than you want,
lower it; if the structure is gone, raise it.
Prerequisites: same as `basic_stable_audio.py`.
"""
from fastvideo import VideoGenerator
PROMPT = "Change the piano to a cello playing the same notes"
# Path to any audio-bearing file (wav, mp3, mp4, m4a, flac, ...).
# Set to `None` to skip A2A and run plain T2A.
INIT_AUDIO_PATH: str | None = None
# Reference fidelity in [0, 1] -- higher = closer to source.
INIT_AUDIO_STRENGTH = 0.6
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/stable-audio-open-1.0-Diffusers",
num_gpus=1,
)
generator.generate_video(
prompt=PROMPT,
output_path="outputs_audio/stable_audio_a2a/output_a2a.wav",
save_video=True,
audio_end_in_s=6.0,
init_audio=INIT_AUDIO_PATH,
init_audio_strength=INIT_AUDIO_STRENGTH,
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,84 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 — inpainting / outpainting (loop extension) example.
User story (loop extension — the killer app):
"I have a 6-second drum loop my client likes. They want it as
background bed for a 30-second ad. I need it to loop seamlessly,
but a hard cut every 6s sounds bad. Let me extend it to 30s,
keeping the first 6s exactly as-is and letting the model continue
the groove for the remaining 24s."
User story (audio repair):
"There's a microphone bump at 0:14 in this 30-second field
recording — really obvious in headphones. Mask out 0:13 to 0:15
and let the model regenerate plausible ambience that blends in.
Everything else stays exactly as I recorded it."
User story (transition smoothing):
"I have two 10-second clips I want to crossfade. Mask out a 1s
overlap region in the middle and let the model invent a coherent
transition between the two."
How it works (RePaint-style blending):
Stable Audio Open 1.0 wasn't trained as an inpainting model
(`model_type=diffusion_cond`, not `diffusion_cond_inpaint`), so we
can't use the upstream's mask-conditioned approach directly. We
use the RePaint trick instead, which works on any v-prediction
diffusion model:
1. Encode the reference clip into latent space.
2. At every denoising step `i`, replace the kept region of the
in-flight latent (where mask == 1) with the reference
re-noised to the next timestep's sigma. Only the unkept
region (mask == 0) is freely denoised.
3. After the loop, the kept region is exactly the reference;
the unkept region is freshly generated content.
This is approximate compared to a properly trained inpainting
checkpoint — the seam between kept/unkept can have slight EQ
discontinuity — but it works on the existing public model.
Tunable: the mask is a 1-D tensor in {0, 1} at the model's sample
rate. Conventions:
1.0 = keep this sample from the reference
0.0 = regenerate this sample
Prerequisites: same as `basic_stable_audio.py`.
"""
import os
from fastvideo import VideoGenerator
PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
# Required: path to the reference audio file (wav, mp3, mp4, m4a, flac,
# ...) you want to extend or repair. The pipeline raises if a mask is
# passed without a reference, so this must be a real path.
REFERENCE_AUDIO_PATH = "path/to/your/loop.wav"
KEEP_SECONDS = 6.0 # first KEEP_SECONDS preserved exactly
TOTAL_SECONDS = 12.0 # extend the loop to this duration
def main() -> None:
if not os.path.isfile(REFERENCE_AUDIO_PATH):
raise FileNotFoundError(
f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
"Edit this script to point at a real audio file (wav/mp3/mp4/"
"m4a/flac) before running.")
generator = VideoGenerator.from_pretrained(
"FastVideo/stable-audio-open-1.0-Diffusers",
num_gpus=1,
)
generator.generate_video(
prompt=PROMPT,
output_path="outputs_audio/stable_audio_inpaint/output_inpaint.wav",
save_video=True,
audio_end_in_s=TOTAL_SECONDS,
inpaint_audio=REFERENCE_AUDIO_PATH,
# Tuple form: keep first KEEP_SECONDS, regenerate the rest.
inpaint_mask=(KEEP_SECONDS, TOTAL_SECONDS),
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,53 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open Small — fast / lightweight T2A example.
User story (interactive UI builder):
"I'm building a sound-design UI where the user types a prompt and
we want sub-2-second feedback so the experience feels like
autocomplete, not a render queue. The full Stable Audio Open 1.0
takes ~8s on a single GPU; the small variant takes a fraction of
that — quality is lower but completely usable for real-time
iteration."
User story (overnight batch jobs):
"I'm generating 10,000 short SFX variants for a procedural game.
Wall-clock matters more than per-clip polish — give me the small
model so I can fit the run in one night instead of a week."
How it works:
The small variant is a separate Stability AI checkpoint
(`stabilityai/stable-audio-open-small`) that ships the same Oobleck
VAE as the 1.0 base model but a smaller / faster DiT (`embed_dim=1024`,
`depth=16`, `qk_norm="ln"`) and only one duration conditioner
(`seconds_total`, no `seconds_start`). FastVideo loads from the
converted Diffusers-format repo `FastVideo/stable-audio-open-small-Diffusers`
via the standard component loader; per-variant arch fields come
from `transformer/config.json` and `conditioner/config.json`.
Prerequisites: same as `basic_stable_audio.py`. The converted repo is
public so no gated-access flow is required.
"""
from fastvideo import VideoGenerator
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/stable-audio-open-small-Diffusers",
num_gpus=1,
)
output_path = "outputs_audio/stable_audio_small/output_stable_audio_small.wav"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
# Small variant trains on a ~11.9s window — keep `audio_end_in_s`
# at or below that.
audio_end_in_s=6.0,
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,73 +0,0 @@
# Cosmos Predict2 2B T2V finetune config.
#
# Data must be preprocessed with Cosmos VAE + T5 text encoder
# into parquet format before training.
models:
student:
_target_: fastvideo.train.models.cosmos.CosmosModel
init_from: nvidia/Cosmos-Predict2-2B-Video2World
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 8
hsdp_shard_dim: 1
data:
data_path: data/cosmos_preprocessed
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
# Cosmos VAE: 4x temporal, 8x spatial compression.
# 93 frames -> 24 latent frames, 480x832 -> 60x104
num_latent_t: 24
num_height: 480
num_width: 832
num_frames: 93
optimizer:
learning_rate: 1.0e-5
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 5000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/cosmos_finetune
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: fastvideo_cosmos
run_name: cosmos_finetune
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos_pipeline.Cosmos2VideoToWorldPipeline
dataset_file: data/cosmos_preprocessed/validation_prompts.json
every_steps: 100
sampling_steps: [50]
guidance_scale: 6.0
pipeline:
flow_shift: 1.0
@@ -1,79 +0,0 @@
# Cosmos-Predict2.5-2B Text-to-World overfitting test config.
#
# Overfits on a few short videos (480x832, 93 frames) to verify the
# Cosmos 2.5 training plugin works end-to-end.
#
# Preprocess data first:
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
#
# Run:
# bash examples/train/run.sh examples/train/configs/overfit_cosmos25_t2w.yaml
models:
student:
_target_: fastvideo.train.models.cosmos.CosmosModel
init_from: KyleShao/Cosmos-Predict2.5-2B-Diffusers
trainable: true
enable_gradient_checkpointing_type: full
flow_shift: 1.0
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 1
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 1
data:
data_path: data/cosmos25_overfit_preprocessed
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 24
num_height: 480
num_width: 832
num_frames: 93
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 300
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/cosmos25_overfit
training_state_checkpointing_steps: 50
checkpoints_total_limit: 2
tracker:
project_name: fastvideo_cosmos25
run_name: cosmos25_overfit
model:
precondition_outputs: false
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos2_5_pipeline.Cosmos2_5Pipeline
dataset_file: data/cosmos25_overfit_preprocessed/validation_prompts.json
every_steps: 150
sampling_steps: [35]
guidance_scale: 7.0
pipeline:
flow_shift: 1.0
@@ -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",
]
-4
View File
@@ -15,7 +15,6 @@ class GenerationResult:
samples: Any | None = None
frames: Any | None = None
audio: Any | None = None
audio_sample_rate: int | None = None
size: tuple[int, int, int] | None = None
generation_time: float | None = None
logging_info: Any | None = None
@@ -45,7 +44,6 @@ class GenerationResult:
"samples",
"frames",
"audio",
"audio_sample_rate",
"size",
"generation_time",
"logging_info",
@@ -64,7 +62,6 @@ class GenerationResult:
samples=result.get("samples"),
frames=result.get("frames"),
audio=result.get("audio"),
audio_sample_rate=result.get("audio_sample_rate"),
size=result.get("size"),
generation_time=result.get("generation_time"),
logging_info=result.get("logging_info"),
@@ -83,7 +80,6 @@ class GenerationResult:
"samples": self.samples,
"frames": self.frames,
"audio": self.audio,
"audio_sample_rate": self.audio_sample_rate,
"size": self.size,
"generation_time": self.generation_time,
"logging_info": self.logging_info,
+6 -48
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,39 +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])
# Stable Audio (T2A): clip start/end in seconds. Honored by
# `StableAudioConditioningStage` + `StableAudioDecodingStage`. Other
# families ignore them.
audio_start_in_s: float | None = None
audio_end_in_s: float | None = None
# Stable Audio audio-to-audio (variation):
# `init_audio` -- a path or `[B, C, samples]` waveform at the model
# sample rate; the pipeline encodes it via the VAE
# and uses it as the starting latent.
# `init_audio_strength` -- 0..1, higher = closer to the reference
# (matches the convention of Stability's
# commercial Stable Audio 2.0 UI). 1.0 ~=
# VAE round-trip, 0.0 ~= plain T2A.
# `init_noise_level` -- legacy raw `sigma_max` override (0.3..500,
# higher = more freedom). Kept for callers
# that already use it; prefer `init_audio_strength`.
init_audio: Any = None
init_audio_strength: float | None = None
init_noise_level: float | None = None
# Stable Audio inpainting (RePaint-style): `inpaint_audio` is the
# reference clip, `inpaint_mask` is a [samples] tensor in {0, 1} where
# 1 means *keep the reference* and 0 means *regenerate*.
inpaint_audio: Any = None
inpaint_mask: Any = None
# 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
@@ -169,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
@@ -185,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
View File
@@ -48,6 +48,8 @@ class ModelConfig:
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(arch_config, key, value)
else:
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
if hasattr(arch_config, "__post_init__"):
arch_config.__post_init__()
+1 -3
View File
@@ -5,13 +5,11 @@ from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"StableAudioConfig"
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig"
]
+1 -2
View File
@@ -50,8 +50,7 @@ class CosmosArchConfig(DiTArchConfig):
})
# Cosmos-specific config parameters based on transformer_cosmos.py
# in_channels includes the condition_mask channel (16 latent + 1 cond = 17)
in_channels: int = 17
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 16
attention_head_dim: int = 128
@@ -1,76 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for the Stable Audio Open 1.0 DiT.
Note: the SA pipeline bypasses the standard `ComposedPipelineBase`
component loader because the published HF repo ships a single monolithic
`model.safetensors` (no Diffusers-style `model_index.json` or
per-subfolder layout). The arch fields and `param_names_mapping` here
document the architecture and key remap so the same conventions used by
the rest of the DiT family apply (FSDP shard conditions, supported
attention backends, future loader integrations) — they are not currently
consumed by `fastvideo/models/loader/fsdp_load.py` for SA.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def _is_transformer_layer(n: str, m) -> bool:
# Matches `transformer.layers.{i}` in the SA DiT module tree.
parts = n.split(".")
return (len(parts) >= 3 and parts[-3] == "transformer" and parts[-2] == "layers" and parts[-1].isdigit())
@dataclass
class StableAudioArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer])
# SA's checkpoint is `stable_audio_tools` raw format (not Diffusers),
# so the only remaps are: strip the `model.model.` host-pipeline
# prefix, and rename `nn.LayerNorm`'s `gamma`/`beta` to torch's
# canonical `weight`/`bias`. Linear / cross-attention naming already
# matches FastVideo's conventions, so no further remap is needed.
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.model\.(.*?)\.gamma$": r"\1.weight",
r"^model\.model\.(.*?)\.beta$": r"\1.bias",
r"^model\.model\.(.*)$": r"\1",
})
# SA only supports backends compatible with single-GPU LocalAttention.
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
# Architecture constants (from the published `model_config.json` for
# `stabilityai/stable-audio-open-1.0`).
io_channels: int = 64
embed_dim: int = 1536
depth: int = 24
num_attention_heads: int = 24
cond_token_dim: int = 768
global_cond_dim: int = 1536
project_cond_tokens: bool = False
project_global_cond: bool = True
# Set to "ln" to wrap attention Q/K in LayerNorm (used by
# `stable-audio-open-small`; absent in the 1.0 base).
qk_norm: str | None = None
def __post_init__(self) -> None:
super().__post_init__()
self.hidden_size = self.embed_dim
self.in_channels = self.io_channels
self.out_channels = self.io_channels
self.num_channels_latents = self.io_channels
self.attention_head_dim = self.embed_dim // self.num_attention_heads
@dataclass
class StableAudioConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=StableAudioArchConfig)
prefix: str = "StableAudio"
@@ -7,12 +7,9 @@ from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
StableAudioConditionerConfig)
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig"
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig"
]
@@ -1,80 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for the Stable Audio Open 1.0 multi-conditioner.
The conditioner bundles three sub-conditioners — a T5 text encoder
(prompt) and two NumberConditioners (`seconds_start` / `seconds_total`)
— into the (cross_attn_cond, cross_attn_mask, global_embed) triple the
DiT consumes. The architecture is fully specified by the official
`stable_audio_tools` `MultiConditioner` config; the constants here
mirror that.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig
from fastvideo.configs.models.encoders.base import (EncoderArchConfig, EncoderConfig)
def _default_configs() -> list[dict]:
"""Default = `stable-audio-open-1.0`'s three sub-conditioners."""
return [
{
"id": "prompt",
"type": "t5",
"config": {
"t5_model_name": "t5-base",
"max_length": 128
}
},
{
"id": "seconds_start",
"type": "number",
"config": {
"min_val": 0,
"max_val": 512
}
},
{
"id": "seconds_total",
"type": "number",
"config": {
"min_val": 0,
"max_val": 512
}
},
]
@dataclass
class StableAudioConditionerArchConfig(EncoderArchConfig):
architectures: list[str] = field(default_factory=lambda: ["StableAudioMultiConditioner"])
# Shared embedding width across all sub-conditioners (T5 last-hidden
# dim and NumberEmbedder feature dim both = `cond_dim`).
cond_dim: int = 768
# Sub-conditioner identifiers. Order in `cross_attention_cond_ids`
# is the concat order for the cross-attn token sequence; order in
# `global_cond_ids` is the concat order for the global FiLM-style
# embedding.
cross_attention_cond_ids: tuple[str, ...] = ("prompt", "seconds_start", "seconds_total")
global_cond_ids: tuple[str, ...] = ("seconds_start", "seconds_total")
# Per-sub-conditioner spec list (mirrors upstream
# `model_config.json.model.conditioning.configs`). Each entry is
# `{"id": ..., "type": "t5"|"number", "config": {...}}`. The default
# matches `stable-audio-open-1.0`; SA-small overrides via the
# `conditioner/config.json` shipped in the converted repo.
configs: list = field(default_factory=_default_configs)
# Match official `stable_audio_tools/models/conditioners.py:334`:
# T5 is loaded directly in fp16.
t5_dtype: str = "float16"
@dataclass
class StableAudioConditionerConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=StableAudioConditionerArchConfig)
prefix: str = "stable_audio_conditioner"
-8
View File
@@ -41,14 +41,6 @@ class T5ArchConfig(TextEncoderArchConfig):
text_len: int = 512
dtype: str | None = None
gradient_checkpointing: bool = False
# Extra fields present in upstream HF T5Config but unused by FastVideo's
# encoder. Declared here so `update_model_arch` doesn't reject them when
# loading repos like `stabilityai/stable-audio-open-1.0` that ship the
# full HF config.
n_positions: int = 512
decoder_start_token_id: int = 0
output_past: bool = True
task_specific_params: dict | None = None
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"),
@@ -5,7 +5,6 @@ from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
from fastvideo.configs.models.vaes.oobleck import OobleckVAEArchConfig, OobleckVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
@@ -17,6 +16,4 @@ __all__ = [
"Gen3CVAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
"OobleckVAEArchConfig",
"OobleckVAEConfig",
]
-68
View File
@@ -1,68 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for the Stable Audio Open 1.0 "Oobleck" VAE.
Mirrors the per-channel `vae/config.json` shipped in
`stabilityai/stable-audio-open-1.0` 1:1 (see
`fastvideo/models/vaes/oobleck.py::OobleckVAE.from_pretrained`, which
constructs the VAE from these fields). Inherits the FastVideo VAEConfig
base so the standard `load_encoder` / `load_decoder` flags + tiling
knobs apply.
Naming: the VAE architecture is officially "Oobleck" (per Stability
AI's stable-audio-tools) — the surrounding model family is "Stable
Audio Open 1.0". This config is named after the architecture
(`OobleckVAEConfig`) since the same VAE is shared across Stable Audio
checkpoints; downstream pipelines reference it by its arch name, not
by a host-pipeline name.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class OobleckVAEArchConfig(VAEArchConfig):
"""Stable Audio Open 1.0 VAE architecture constants."""
architectures: list[str] = field(default_factory=lambda: ["AutoencoderOobleck"])
# From stabilityai/stable-audio-open-1.0/vae/config.json.
encoder_hidden_size: int = 128
downsampling_ratios: list[int] = field(default_factory=lambda: [2, 4, 4, 8, 8])
channel_multiples: list[int] = field(default_factory=lambda: [1, 2, 4, 8, 16])
decoder_channels: int = 128
decoder_input_channels: int = 64
audio_channels: int = 2 # stereo
sampling_rate: int = 44100
@dataclass
class OobleckVAEConfig(VAEConfig):
"""FastVideo VAE config wrapping the Oobleck arch.
Audio VAEs don't use the temporal/spatial tiling defaults that the
base VAEConfig is shaped for (those exist for video VAEs); they are
retained but irrelevant for audio.
"""
arch_config: VAEArchConfig = field(default_factory=OobleckVAEArchConfig)
# Audio is 1-D, so the video-VAE tiling defaults are inert. Disable
# them so callers don't accidentally trip on tile-stride math built
# for spatial tensors.
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
# Where the FastVideo loader / pipeline-glue wrapper should fetch
# weights from when no local path is supplied. Gated repo — caller's
# HF token must have accepted terms on
# https://huggingface.co/stabilityai/stable-audio-open-1.0.
pretrained_path: str = "stabilityai/stable-audio-open-1.0"
pretrained_subfolder: str = "vae"
# Match official `stable_audio_tools`: VAE runs in fp16 (the
# `pretransform.model_half` path in
# `stable_audio_tools/models/pretransforms.py`).
pretrained_dtype: str = "float16"
+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, ))
@@ -1,67 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""`PipelineConfig` for Stable Audio Open 1.0."""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, VAEConfig
from fastvideo.configs.models.dits import StableAudioConfig
from fastvideo.configs.models.vaes import OobleckVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class StableAudioT2AConfig(PipelineConfig):
"""Stable Audio Open 1.0 pipeline config."""
dit_config: DiTConfig = field(default_factory=StableAudioConfig)
# Standard `TransformerLoader` reads `dit_precision`; default in
# `PipelineConfig` is bf16, but we want fp16 to match official.
dit_precision: str = "fp16"
vae_config: VAEConfig = field(default_factory=OobleckVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# `StableAudioMultiConditioner` owns its own T5; zero out the
# parent's text-encoder slots so the length-equality validator passes.
text_encoder_configs: tuple = field(default_factory=tuple)
preprocess_text_funcs: tuple = field(default_factory=tuple)
postprocess_text_funcs: tuple = field(default_factory=tuple)
num_inference_steps: int = 100
guidance_scale: float = 7.0
audio_end_in_s: float = 10.0 # short-clip default
audio_start_in_s: float = 0.0
sampling_rate: int = 44100
audio_channels: int = 2
# Stable Audio Open 1.0 was trained at a fixed 2,097,152-sample
# window (= 2097152 / 44100 ≈ 47.55s). Anything past this is
# silently truncated by the post-decode slice — validate up-front.
sample_size: int = 2097152
max_audio_duration_s: float = 2097152 / 44100
# Match the official `stable_audio_tools` defaults (`model_half=True`
# in `run_gradio.py`), which loads the DiT, VAE, and T5 in fp16 and
# wraps T5 forward in `autocast(fp16)`. fp16 is also a hard
# requirement for FlashAttention-2 / FA-3.
precision: str = "fp16"
vae_precision: str = "fp16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
def __post_init__(self) -> None:
# A2A needs encode; load both halves for either path.
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class StableAudioOpenSmallConfig(StableAudioT2AConfig):
"""`stable-audio-open-small` overrides: shorter training window
(524288 samples ≈ 11.89s @ 44.1 kHz) and a faster default sampler
config carried by the small preset.
"""
sample_size: int = 524288
max_audio_duration_s: float = 524288 / 44100
audio_end_in_s: float = 6.0 # short-clip default suitable for the small window
+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",
]
+36 -78
View File
@@ -627,32 +627,18 @@ class VideoGenerator:
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs (skip the make_grid loop for audio-only, where
# `samples` is a 1×3×1×8×8 placeholder no caller will use).
audio_only = bool(output_batch.extra.get("audio_only"))
frames: list[np.ndarray] = []
if not audio_only:
videos = rearrange(samples, "b c t h w -> t b c h w")
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.permute(1, 2, 0).squeeze(-1)
x = (x * 255).to(torch.uint8)
frames.append(x.cpu().numpy())
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.permute(1, 2, 0).squeeze(-1)
x = (x * 255).to(torch.uint8)
frames.append(x.cpu().numpy())
# Save output if requested
if batch.save_video:
if output_batch.extra.get("audio_only"):
# Audio-only workload: write a standalone .wav rather than
# muxing the audio into a placeholder mp4 (which forces
# ffmpeg to round 8x8 placeholder frames up to 16x16).
output_path = self._rewrite_extension(output_path, ".wav")
self._write_pcm_wav(
output_path,
output_batch.extra["audio"],
int(output_batch.extra["audio_sample_rate"]),
)
logger.info("Saved audio to %s", output_path)
elif self._is_image_workload():
if self._is_image_workload():
# Image workloads (t2i, i2i, …): save the first frame as PNG.
imageio.imwrite(output_path, frames[0])
logger.info("Saved image to %s", output_path)
@@ -669,11 +655,7 @@ class VideoGenerator:
"prompts": prompt,
"samples": samples if batch.return_frames else None,
"frames": frames if batch.return_frames else None,
# Audio is the primary output for audio workloads — return it
# whenever the pipeline produced one, regardless of
# `return_frames` (which gates the video-shaped buffers).
"audio": output_batch.extra.get("audio"),
"audio_sample_rate": output_batch.extra.get("audio_sample_rate"),
"audio": output_batch.extra.get("audio") if batch.return_frames else None,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
@@ -701,55 +683,7 @@ class VideoGenerator:
return result.to_legacy_dict()
@staticmethod
def _rewrite_extension(path: str, new_ext: str) -> str:
root, old_ext = os.path.splitext(path)
new_path = root + new_ext
if old_ext and old_ext.lower() != new_ext.lower():
logger.info("Rewriting output extension %s -> %s.", old_ext, new_ext)
return new_path
@staticmethod
def _audio_to_int16(audio: torch.Tensor | np.ndarray, ) -> tuple[np.ndarray, int]:
"""Normalize `[samples]` / `[samples, channels]` / `[channels,
samples]` audio in roughly [-1, 1] to a `(int16 [samples,
channels], num_channels)` pair. Raises `ValueError` for shapes
we can't classify.
"""
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
raise ValueError(f"Unexpected audio shape {audio_np.shape}.")
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
return audio_int16, audio_int16.shape[1]
@classmethod
def _write_pcm_wav(
cls,
wav_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
) -> int:
"""Write 16-bit PCM WAV; returns the channel count."""
import wave
audio_int16, num_channels = cls._audio_to_int16(audio)
with wave.open(wav_path, "wb") as f:
f.setnchannels(num_channels)
f.setsampwidth(2)
f.setframerate(sample_rate)
f.writeframes(audio_int16.tobytes())
return num_channels
@classmethod
def _mux_audio(
cls,
video_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
@@ -762,13 +696,37 @@ class VideoGenerator:
"Install with: pip install av")
return False
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
logger.warning("Unexpected audio shape %s; skipping mux.", audio_np.shape)
return False
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
num_channels = audio_int16.shape[1]
layout = "stereo" if num_channels == 2 else "mono"
try:
import wave
with tempfile.TemporaryDirectory() as tmpdir:
out_path = os.path.join(tmpdir, "muxed.mp4")
wav_path = os.path.join(tmpdir, "audio.wav")
num_channels = cls._write_pcm_wav(wav_path, audio, sample_rate)
layout = "stereo" if num_channels == 2 else "mono"
# Write audio to WAV file
with wave.open(wav_path, "wb") as wav_file:
wav_file.setnchannels(num_channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(audio_int16.tobytes())
# Open input video and audio
input_video = av.open(video_path)
+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 -35
View File
@@ -556,10 +556,7 @@ class CosmosTransformer3DModel(BaseDiT):
self.extra_pos_embed_type = config.extra_pos_embed_type
# 1. Patch Embedding
# config.in_channels already includes the condition_mask channel
# (HF config: in_channels=17 = 16 latent + 1 condition_mask).
# Only add +1 for the padding_mask when concat_padding_mask=True.
patch_embed_in_channels = config.in_channels + (1 if config.concat_padding_mask else 0)
patch_embed_in_channels = config.in_channels + 1 if config.concat_padding_mask else config.in_channels
self.patch_embed = CosmosPatchEmbed(patch_embed_in_channels,
inner_dim,
config.patch_size,
@@ -620,28 +617,6 @@ class CosmosTransformer3DModel(BaseDiT):
batch_size, num_channels, num_frames, height, width = hidden_states.shape
# Defensive dtype alignment: the Cosmos checkpoint is bf16 but
# FSDP-wrapped training copies may report fp32 via
# `next(parameters()).dtype`, which disables autocast in the
# shared denoising stage and feeds fp32 tensors into bf16
# weights. Cast every external input to the patch_embed weight
# dtype so the model forward is robust regardless of caller.
_target_dtype = self.patch_embed.proj.weight.dtype
if hidden_states.dtype != _target_dtype:
hidden_states = hidden_states.to(_target_dtype)
if condition_mask is not None and condition_mask.dtype != _target_dtype:
condition_mask = condition_mask.to(_target_dtype)
if padding_mask is not None and padding_mask.dtype != _target_dtype:
padding_mask = padding_mask.to(_target_dtype)
if isinstance(encoder_hidden_states, torch.Tensor):
if encoder_hidden_states.dtype != _target_dtype:
encoder_hidden_states = encoder_hidden_states.to(_target_dtype)
else:
encoder_hidden_states = [
t.to(_target_dtype) if t.dtype != _target_dtype else t
for t in encoder_hidden_states
]
# 1. Concatenate padding mask if needed & prepare attention mask
if condition_mask is not None:
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
@@ -651,10 +626,6 @@ class CosmosTransformer3DModel(BaseDiT):
padding_mask = transforms.functional.resize(
padding_mask, list(hidden_states.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST
)
# torchvision.resize may upcast bf16/fp16 → fp32; restore
# hidden_states' dtype so the subsequent cat doesn't promote
# everything and break patch_embed (bf16 weights).
padding_mask = padding_mask.to(hidden_states.dtype)
hidden_states = torch.cat(
[hidden_states, padding_mask.unsqueeze(2).repeat(batch_size, 1, num_frames, 1, 1)], dim=1
)
@@ -732,11 +703,6 @@ class CosmosTransformer3DModel(BaseDiT):
hidden_states = hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5)
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
# Return as tuple for compatibility with callers that
# do `transformer(..., return_dict=False)[0]` (diffusers
# convention used by CosmosDenoisingStage).
if not kwargs.get("return_dict", True):
return (hidden_states,)
return hidden_states
# Entry point for model registry
-389
View File
@@ -1,389 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 DiT.
Continuous transformer with rotary self-attention, GQA cross-attention,
and prepend global conditioning. 24 layers, embed_dim=1536, head_dim=64.
"""
from __future__ import annotations
import math
from typing import Any
import torch
from einops import rearrange
from torch import nn
from fastvideo.attention import LocalAttention
from fastvideo.configs.models.dits import StableAudioConfig
from fastvideo.layers.layernorm import FP32LayerNorm
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.loader.utils import get_param_names_mapping
# Single import-time snapshot — re-reading via `StableAudioConfig()` per
# `Attention.__init__` would rebuild the nested dataclass + regex map ~48
# times during a single DiT construction. Reused for the class-level
# attribute defaults below.
_DEFAULT_CONFIG = StableAudioConfig()
_SUPPORTED_BACKENDS = _DEFAULT_CONFIG.arch_config._supported_attention_backends
class FourierFeatures(nn.Module):
"""Random-Fourier learned-frequency timestep encoder."""
def __init__(self, in_features: int, out_features: int, std: float = 1.0) -> None:
super().__init__()
assert out_features % 2 == 0
self.weight = nn.Parameter(torch.randn([out_features // 2, in_features]) * std)
def forward(self, x: torch.Tensor) -> torch.Tensor:
f = 2 * math.pi * x @ self.weight.T
return torch.cat([f.cos(), f.sin()], dim=-1)
# Partial-rotary with halves-swap (`unbind(-2)`, `[-x2, x1]`). Different
# from FastVideo's `_apply_rotary_emb` (interleaved pairs, `unbind(-1)`),
# so kept local.
class RotaryEmbedding(nn.Module):
def __init__(self, dim: int, base: float = 10000.0) -> None:
super().__init__()
inv_freq = 1.0 / (base**(torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq)
self.register_buffer("scale", None)
def forward_from_seq_len(self, seq_len: int):
t = torch.arange(seq_len, device=self.inv_freq.device, dtype=torch.float32)
freqs = torch.einsum("i , j -> i j", t, self.inv_freq)
freqs = torch.cat((freqs, freqs), dim=-1)
return freqs, 1.0
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
x = rearrange(x, "... (j d) -> ... j d", j=2)
x1, x2 = x.unbind(dim=-2)
return torch.cat((-x2, x1), dim=-1)
def _apply_rotary_pos_emb(t: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
out_dtype = t.dtype
rot_dim, seq_len = freqs.shape[-1], t.shape[-2]
freqs = freqs.to(torch.float32)[-seq_len:, :]
t = t.to(torch.float32)
if t.ndim == 4 and freqs.ndim == 3:
freqs = rearrange(freqs, "b n d -> b 1 n d")
t_rot, t_unrot = t[..., :rot_dim], t[..., rot_dim:]
t_rot = (t_rot * freqs.cos()) + (_rotate_half(t_rot) * freqs.sin())
return torch.cat((t_rot.to(out_dtype), t_unrot.to(out_dtype)), dim=-1)
# SwiGLU FF — local because `fastvideo.layers.mlp.MLP` is non-gated.
class _GLU(nn.Module):
def __init__(self, dim_in: int, dim_out: int, activation: nn.Module) -> None:
super().__init__()
self.act = activation
self.proj = ReplicatedLinear(dim_in, dim_out * 2, bias=True)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x, _ = self.proj(x)
x, gate = x.chunk(2, dim=-1)
return x * self.act(gate)
class FeedForward(nn.Module):
# Sequential layout `(GLU, Identity, Linear, Identity)` keeps the
# checkpoint keys at indices 0 and 2.
def __init__(self, dim: int, mult: int = 4, zero_init_output: bool = True) -> None:
super().__init__()
inner_dim = int(dim * mult)
linear_in = _GLU(dim, inner_dim, nn.SiLU())
linear_out = ReplicatedLinear(inner_dim, dim, bias=True)
if zero_init_output:
nn.init.zeros_(linear_out.weight)
nn.init.zeros_(linear_out.bias)
self.ff = nn.Sequential(linear_in, nn.Identity(), linear_out, nn.Identity())
def forward(self, x: torch.Tensor) -> torch.Tensor:
for mod in self.ff:
if isinstance(mod, ReplicatedLinear):
x, _ = mod(x)
else:
x = mod(x)
return x
# Cross-attention is GQA (24 query heads, 12 KV heads); both backends
# (FlashAttn, SDPA with `enable_gqa=True`) handle it.
class Attention(nn.Module):
def __init__(self, dim: int, dim_heads: int = 64, dim_context: int | None = None,
zero_init_output: bool = True, qk_norm: str | None = None) -> None:
super().__init__()
self.dim = dim
self.dim_heads = dim_heads
dim_kv = dim_context if dim_context is not None else dim
self.num_heads = dim // dim_heads
self.kv_heads = dim_kv // dim_heads
if dim_context is not None:
self.to_q = ReplicatedLinear(dim, dim, bias=False)
self.to_kv = ReplicatedLinear(dim_kv, dim_kv * 2, bias=False)
else:
self.to_qkv = ReplicatedLinear(dim, dim * 3, bias=False)
self.to_out = ReplicatedLinear(dim, dim, bias=False)
if zero_init_output:
nn.init.zeros_(self.to_out.weight)
# `stable-audio-open-small` wraps Q/K in LayerNorm before attn
# (`attn_kwargs.qk_norm = "ln"` in its `model_config.json`); the
# 1.0 base does not. Names match upstream (`q_norm`/`k_norm`)
# so the converted state dict loads strict.
if qk_norm == "ln":
self.q_norm = nn.LayerNorm(dim_heads)
self.k_norm = nn.LayerNorm(dim_heads)
elif qk_norm is None:
self.q_norm = nn.Identity()
self.k_norm = nn.Identity()
else:
raise ValueError(f"Unsupported qk_norm={qk_norm!r}; expected 'ln' or None.")
self.attn = LocalAttention(num_heads=self.num_heads, head_size=dim_heads,
num_kv_heads=self.kv_heads, causal=False,
supported_attention_backends=_SUPPORTED_BACKENDS)
def forward(self, x: torch.Tensor, context: torch.Tensor | None = None,
rotary_pos_emb: tuple[torch.Tensor, float] | None = None) -> torch.Tensor:
h, kv_h, has_context = self.num_heads, self.kv_heads, context is not None
kv_input = context if has_context else x
if has_context:
q, _ = self.to_q(x)
kv, _ = self.to_kv(kv_input)
k, v = kv.chunk(2, dim=-1)
else:
qkv, _ = self.to_qkv(x)
q, k, v = qkv.chunk(3, dim=-1)
# LocalAttention expects [batch, seq_len, num_heads, head_dim].
q = rearrange(q, "b n (h d) -> b n h d", h=h)
k = rearrange(k, "b n (h d) -> b n h d", h=kv_h)
v = rearrange(v, "b n (h d) -> b n h d", h=kv_h)
q = self.q_norm(q)
k = self.k_norm(k)
if rotary_pos_emb is not None:
freqs, _ = rotary_pos_emb
v_dtype = v.dtype
# Partial rotary (rot_dim < head_dim) with halves-swap, so
# apply outside LocalAttention. q,k come in as [B, S, H, D];
# transpose to [B, H, S, D] for the helper.
q_t = q.transpose(1, 2)
k_t = k.transpose(1, 2)
if q_t.shape[-2] >= k_t.shape[-2]:
ratio = q_t.shape[-2] / k_t.shape[-2]
q_freqs, k_freqs = freqs, ratio * freqs
else:
ratio = k_t.shape[-2] / q_t.shape[-2]
q_freqs, k_freqs = ratio * freqs, freqs
q = _apply_rotary_pos_emb(q_t, q_freqs).to(v_dtype).transpose(1, 2)
k = _apply_rotary_pos_emb(k_t, k_freqs).to(v_dtype).transpose(1, 2)
out = self.attn(q, k, v)
out = rearrange(out, "b n h d -> b n (h d)")
out, _ = self.to_out(out)
return out
class TransformerBlock(nn.Module):
def __init__(self, dim: int, dim_heads: int = 64, cross_attend: bool = False,
dim_context: int | None = None, zero_init_branch_outputs: bool = True,
qk_norm: str | None = None) -> None:
super().__init__()
self.dim = dim
self.dim_heads = min(dim_heads, dim)
self.cross_attend = cross_attend
self.pre_norm = FP32LayerNorm(dim, elementwise_affine=True)
self.self_attn = Attention(dim, dim_heads=self.dim_heads,
zero_init_output=zero_init_branch_outputs,
qk_norm=qk_norm)
if cross_attend:
self.cross_attend_norm = FP32LayerNorm(dim, elementwise_affine=True)
self.cross_attn = Attention(dim, dim_heads=self.dim_heads, dim_context=dim_context,
zero_init_output=zero_init_branch_outputs,
qk_norm=qk_norm)
self.ff_norm = FP32LayerNorm(dim, elementwise_affine=True)
self.ff = FeedForward(dim, zero_init_output=zero_init_branch_outputs)
def forward(self, x: torch.Tensor, context: torch.Tensor | None = None,
rotary_pos_emb: tuple[torch.Tensor, float] | None = None) -> torch.Tensor:
x = x + self.self_attn(self.pre_norm(x), rotary_pos_emb=rotary_pos_emb)
if context is not None and self.cross_attend:
x = x + self.cross_attn(self.cross_attend_norm(x), context=context)
x = x + self.ff(self.ff_norm(x))
return x
class ContinuousTransformer(nn.Module):
def __init__(self, dim: int, depth: int, *, dim_heads: int = 64, dim_in: int | None = None,
dim_out: int | None = None, cross_attend: bool = False,
cond_token_dim: int | None = None, zero_init_branch_outputs: bool = True,
qk_norm: str | None = None) -> None:
super().__init__()
self.dim = dim
self.depth = depth
self.project_in = (ReplicatedLinear(dim_in, dim, bias=False) if dim_in is not None
else nn.Identity())
self.project_out = (ReplicatedLinear(dim, dim_out, bias=False) if dim_out is not None
else nn.Identity())
self.rotary_pos_emb = RotaryEmbedding(max(dim_heads // 2, 32))
self.layers = nn.ModuleList([
TransformerBlock(dim, dim_heads=dim_heads, cross_attend=cross_attend,
dim_context=cond_token_dim,
zero_init_branch_outputs=zero_init_branch_outputs,
qk_norm=qk_norm) for _ in range(depth)
])
def forward(self, x: torch.Tensor, prepend_embeds: torch.Tensor | None = None,
context: torch.Tensor | None = None) -> torch.Tensor:
if isinstance(self.project_in, ReplicatedLinear):
x, _ = self.project_in(x)
if prepend_embeds is not None:
assert prepend_embeds.shape[-1] == x.shape[-1]
x = torch.cat((prepend_embeds, x), dim=-2)
rotary = self.rotary_pos_emb.forward_from_seq_len(x.shape[1])
for layer in self.layers:
x = layer(x, context=context, rotary_pos_emb=rotary)
if isinstance(self.project_out, ReplicatedLinear):
x, _ = self.project_out(x)
return x
class StableAudioDiT(BaseDiT):
"""Stable Audio Open 1.0 diffusion transformer."""
_fsdp_shard_conditions = _DEFAULT_CONFIG.arch_config._fsdp_shard_conditions
_compile_conditions = _DEFAULT_CONFIG.arch_config._compile_conditions
param_names_mapping = _DEFAULT_CONFIG.arch_config.param_names_mapping
reverse_param_names_mapping: dict = {}
def __init__(self, config: StableAudioConfig | None = None,
hf_config: dict[str, Any] | None = None) -> None:
if config is None:
config = StableAudioConfig()
super().__init__(config=config, hf_config=hf_config or {})
arch = config.arch_config
self.hidden_size = arch.hidden_size
self.num_attention_heads = arch.num_attention_heads
self.num_channels_latents = arch.num_channels_latents
io_channels = arch.io_channels
embed_dim = arch.embed_dim
depth = arch.depth
num_heads = arch.num_attention_heads
cond_token_dim = arch.cond_token_dim
global_cond_dim = arch.global_cond_dim
project_cond_tokens = arch.project_cond_tokens
project_global_cond = arch.project_global_cond
qk_norm = arch.qk_norm
self.cond_token_dim = cond_token_dim
timestep_features_dim = 256
self.timestep_features = FourierFeatures(1, timestep_features_dim)
self.to_timestep_embed = nn.Sequential(
ReplicatedLinear(timestep_features_dim, embed_dim, bias=True),
nn.SiLU(),
ReplicatedLinear(embed_dim, embed_dim, bias=True),
)
self.diffusion_objective = "v"
cond_embed_dim = cond_token_dim if not project_cond_tokens else embed_dim
self.to_cond_embed = nn.Sequential(
ReplicatedLinear(cond_token_dim, cond_embed_dim, bias=False),
nn.SiLU(),
ReplicatedLinear(cond_embed_dim, cond_embed_dim, bias=False),
)
global_embed_dim = global_cond_dim if not project_global_cond else embed_dim
self.to_global_embed = nn.Sequential(
ReplicatedLinear(global_cond_dim, global_embed_dim, bias=False),
nn.SiLU(),
ReplicatedLinear(global_embed_dim, global_embed_dim, bias=False),
)
self.transformer = ContinuousTransformer(
dim=embed_dim, depth=depth, dim_heads=embed_dim // num_heads, dim_in=io_channels,
dim_out=io_channels, cross_attend=True, cond_token_dim=cond_embed_dim,
qk_norm=qk_norm,
)
self.preprocess_conv = nn.Conv1d(io_channels, io_channels, 1, bias=False)
nn.init.zeros_(self.preprocess_conv.weight)
self.postprocess_conv = nn.Conv1d(io_channels, io_channels, 1, bias=False)
nn.init.zeros_(self.postprocess_conv.weight)
self.io_channels = io_channels
self.embed_dim = embed_dim
self.depth = depth
self.num_heads = num_heads
self.__post_init__()
@staticmethod
def _seq_apply(seq: nn.Sequential, x: torch.Tensor) -> torch.Tensor:
for mod in seq:
if isinstance(mod, ReplicatedLinear):
x, _ = mod(x)
else:
x = mod(x)
return x
def forward(self, x: torch.Tensor, t: torch.Tensor, *, cross_attn_cond: torch.Tensor,
global_embed: torch.Tensor) -> torch.Tensor:
"""Forward over a single batch. CFG batching is the caller's job."""
model_dtype = next(self.parameters()).dtype
x = x.to(model_dtype)
t = t.to(model_dtype)
cross_attn_cond = cross_attn_cond.to(model_dtype)
global_embed = global_embed.to(model_dtype)
cross_attn_cond = self._seq_apply(self.to_cond_embed, cross_attn_cond)
global_embed = self._seq_apply(self.to_global_embed, global_embed)
timestep_embed = self._seq_apply(self.to_timestep_embed, self.timestep_features(t[:, None]))
global_embed = global_embed + timestep_embed
prepend_inputs = global_embed.unsqueeze(1)
x = self.preprocess_conv(x) + x
x = rearrange(x, "b c t -> b t c")
out = self.transformer(x, prepend_embeds=prepend_inputs, context=cross_attn_cond)
out = rearrange(out, "b t c -> b c t")[:, :, prepend_inputs.shape[1]:]
return self.postprocess_conv(out) + out
@classmethod
def from_official_state_dict(cls, state_dict: dict[str, torch.Tensor],
prefix: str = "model.model.") -> "StableAudioDiT":
"""Load from a raw `stable_audio_tools` monolithic state dict.
Kept for tests / older checkpoints; production loads go through
the standard `TransformerLoader` against the converted Diffusers
repo.
"""
model = cls()
mapping_fn = get_param_names_mapping(model.config.arch_config.param_names_mapping)
remapped: dict[str, torch.Tensor] = {}
for k, v in state_dict.items():
if not k.startswith(prefix):
continue
new_key, _, _ = mapping_fn(k)
remapped[new_key] = v
missing, unexpected = model.load_state_dict(remapped, strict=True)
if missing or unexpected:
raise RuntimeError(
f"StableAudioDiT load mismatch — missing={missing[:5]} unexpected={unexpected[:5]}")
return model
EntryClass = StableAudioDiT
@@ -1,214 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 conditioner.
T5-base text encoder + two NumberConditioners (`seconds_start`,
`seconds_total`), wrapped by `StableAudioMultiConditioner` which
produces the cross-attention and global-conditioning tensors the DiT
expects.
"""
from __future__ import annotations
import math
import torch
import torch.nn as nn
from einops import rearrange
from fastvideo.configs.models.encoders import StableAudioConditionerConfig
class _LearnedPositionalEmbedding(nn.Module):
def __init__(self, dim: int) -> None:
super().__init__()
assert (dim % 2) == 0
self.weights = nn.Parameter(torch.randn(dim // 2))
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(x, "b -> b 1")
freqs = x * rearrange(self.weights, "d -> 1 d") * 2 * math.pi
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim=-1)
return torch.cat((x, fouriered), dim=-1)
def _time_positional_embedding(dim: int, out_features: int) -> nn.Sequential:
return nn.Sequential(_LearnedPositionalEmbedding(dim),
nn.Linear(in_features=dim + 1, out_features=out_features))
class NumberEmbedder(nn.Module):
def __init__(self, features: int, dim: int = 256) -> None:
super().__init__()
self.features = features
self.embedding = _time_positional_embedding(dim=dim, out_features=features)
def forward(self, x: torch.Tensor | list[float]) -> torch.Tensor:
if not torch.is_tensor(x):
device = next(self.embedding.parameters()).device
x = torch.tensor(x, device=device)
shape = x.shape
x = rearrange(x, "... -> (...)")
out = self.embedding(x)
return out.view(*shape, self.features)
class _Conditioner(nn.Module):
def __init__(self, dim: int, output_dim: int, project_out: bool = False) -> None:
super().__init__()
self.dim = dim
self.output_dim = output_dim
self.proj_out = (nn.Linear(dim, output_dim) if dim != output_dim or project_out
else nn.Identity())
class T5Conditioner(_Conditioner):
"""T5 text conditioner. Pads to `model_max_length` (=128 for the SA
repo's tokenizer, NOT the standard 512) and emits a masked
last-hidden-state.
"""
T5_MODEL_DIMS = {"t5-base": 768}
def __init__(self, output_dim: int, t5_model_name: str = "t5-base",
max_length: int = 128, dtype: str = "float16") -> None:
super().__init__(self.T5_MODEL_DIMS[t5_model_name], output_dim, project_out=False)
from transformers import AutoTokenizer, T5EncoderModel
self.max_length = max_length
self.tokenizer = AutoTokenizer.from_pretrained(t5_model_name)
# T5 loaded directly in fp16 (config-driven) to match official
# `stable_audio_tools/models/conditioners.py:334`. Registered as
# a normal submodule so `.to(device)` / `torch.compile` track it;
# `from_official_state_dict` filters `conditioners.prompt.*` from
# the missing-key check (T5 weights are absent from the SA
# checkpoint by design).
# Explicit lookup so a typo (e.g. "fp16" instead of "float16") errors
# at load time rather than silently falling back to a wrong dtype.
torch_dtype = getattr(torch, dtype)
if not isinstance(torch_dtype, torch.dtype):
raise ValueError(f"T5Conditioner dtype={dtype!r} is not a torch.dtype.")
self._t5_dtype = torch_dtype
self.model = (T5EncoderModel.from_pretrained(t5_model_name).eval().requires_grad_(False).to(torch_dtype))
def forward(self, texts: list[str], device: torch.device | str) -> tuple[torch.Tensor, torch.Tensor]:
encoded = self.tokenizer(texts, truncation=True, max_length=self.max_length,
padding="max_length", return_tensors="pt")
input_ids = encoded["input_ids"].to(device)
attention_mask = encoded["attention_mask"].to(device).to(torch.bool)
# Mirror official's `autocast(fp16)` wrap on T5 forward.
with torch.no_grad(), torch.autocast(device_type="cuda", dtype=self._t5_dtype):
embeddings = self.model(input_ids=input_ids,
attention_mask=attention_mask)["last_hidden_state"]
embeddings = self.proj_out(embeddings) * attention_mask.unsqueeze(-1).float()
return embeddings, attention_mask
class NumberConditioner(_Conditioner):
"""Float-valued conditioner with min/max clamping + NumberEmbedder."""
def __init__(self, output_dim: int, min_val: float = 0, max_val: float = 1) -> None:
super().__init__(output_dim, output_dim)
self.min_val = min_val
self.max_val = max_val
self.embedder = NumberEmbedder(features=output_dim)
def forward(self, floats: list[float], device: torch.device | str) -> tuple[torch.Tensor, torch.Tensor]:
floats = [float(x) for x in floats]
floats_t = torch.tensor(floats, device=device).clamp(self.min_val, self.max_val)
normalized = (floats_t - self.min_val) / (self.max_val - self.min_val)
emb_dtype = next(self.embedder.parameters()).dtype
normalized = normalized.to(emb_dtype)
float_embeds = self.embedder(normalized).unsqueeze(1)
return float_embeds, torch.ones(float_embeds.shape[0], 1, device=device)
class StableAudioMultiConditioner(nn.Module):
"""SA-Open-1.0 conditioner: T5 prompt + duration NumberConditioners.
All hardcoded constants (cond_dim, sub-conditioner ids, T5 model
name + max_length, NumberConditioner ranges) live on
`StableAudioConditionerConfig` — see
`fastvideo/configs/models/encoders/stable_audio_conditioner.py`.
"""
def __init__(self, config: StableAudioConditionerConfig | None = None) -> None:
super().__init__()
self.config = config or StableAudioConditionerConfig()
arch = self.config.arch_config
# Build sub-conditioners from the `configs` list (mirrors
# upstream's `MultiConditioner` factory).
sub: dict[str, nn.Module] = {}
for spec in arch.configs:
sid = spec["id"]
stype = spec["type"]
scfg = spec["config"]
if stype == "t5":
sub[sid] = T5Conditioner(output_dim=arch.cond_dim,
t5_model_name=scfg["t5_model_name"],
max_length=scfg["max_length"],
dtype=arch.t5_dtype)
elif stype == "number":
sub[sid] = NumberConditioner(output_dim=arch.cond_dim,
min_val=scfg["min_val"], max_val=scfg["max_val"])
else:
raise ValueError(f"Unknown sub-conditioner type {stype!r} for id {sid!r}.")
self.conditioners = nn.ModuleDict(sub)
self.cross_attention_cond_ids = tuple(arch.cross_attention_cond_ids)
self.global_cond_ids = tuple(arch.global_cond_ids)
def forward(self, batch_metadata: list[dict],
device: torch.device | str) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
out: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
for key, conditioner in self.conditioners.items():
inputs = [x[key] for x in batch_metadata]
out[key] = conditioner(inputs, device)
return out
def get_conditioning_inputs(
self, cond: dict[str, tuple[torch.Tensor, torch.Tensor]]
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Pack conditioner outputs into the (cross_attn_cond,
cross_attn_mask, global_embed) triple the DiT consumes. Order
is driven by `cross_attention_cond_ids` / `global_cond_ids`
from the config — SA-1.0 uses three sub-conditioners
(prompt + seconds_start + seconds_total); SA-small uses two
(prompt + seconds_total).
"""
x_embs = [cond[i][0] for i in self.cross_attention_cond_ids]
x_masks = [cond[i][1] for i in self.cross_attention_cond_ids]
cross_attn_cond = torch.cat(x_embs, dim=1)
cross_attn_mask = torch.cat(x_masks, dim=1)
global_embed = torch.cat([cond[i][0][:, 0] for i in self.global_cond_ids], dim=-1)
return cross_attn_cond, cross_attn_mask, global_embed
@classmethod
def from_official_state_dict(cls, state_dict: dict[str, torch.Tensor],
prefix: str = "conditioner.") -> "StableAudioMultiConditioner":
"""Load NumberConditioner weights from a raw `stable_audio_tools`
monolithic state dict. Kept for tests / older checkpoints;
production loads go through the standard `ConditionerLoader`
against the converted Diffusers repo.
"""
mc = cls()
own_state = mc.state_dict()
loaded: dict[str, torch.Tensor] = {}
for k, v in state_dict.items():
if not k.startswith(prefix):
continue
stripped = k[len(prefix):]
if stripped in own_state:
loaded[stripped] = v
# T5 keys are intentionally absent from the checkpoint.
missing = [k for k in own_state.keys() if k not in loaded
and not k.startswith("conditioners.prompt.")]
unexpected = [k for k in loaded.keys() if k not in own_state]
if missing or unexpected:
raise RuntimeError(
f"StableAudioMultiConditioner load mismatch — missing={missing[:5]} unexpected={unexpected[:5]}"
)
mc.load_state_dict(loaded, strict=False)
return mc
EntryClass = StableAudioMultiConditioner
+224
View File
@@ -0,0 +1,224 @@
# SPDX-License-Identifier: Apache-2.0
# Inspired by SGLang's layerwise offload implementation:
# https://github.com/sgl-project/sglang/pull/15511
#
# This implementation provides a lightweight layerwise CPU offload manager
# with async H2D prefetch using a dedicated CUDA stream, following SGLang's design.
import re
from contextlib import contextmanager
from typing import Dict, Set, Optional, Tuple
import torch
class LayerwiseOffloadManager:
"""A lightweight layerwise CPU offload manager.
Offloads per-layer parameters/buffers from GPU to CPU, and supports async H2D
prefetch using a dedicated CUDA stream.
"""
def __init__(
self,
model: torch.nn.Module,
*,
module_list_attr: str,
num_layers: int,
enabled: bool,
pin_cpu_memory: bool = True,
auto_initialize: bool = False,
) -> None:
self.model = model
self.module_list_attr = module_list_attr
self.num_layers = int(num_layers)
self.pin_cpu_memory = bool(pin_cpu_memory)
self.enabled = bool(enabled and torch.cuda.is_available())
self.device = (
torch.device("cuda", torch.cuda.current_device()) if self.enabled else None
)
self.copy_stream = torch.cuda.Stream() if self.enabled else None
self._layer_name_re = re.compile(
rf"(^|\.){re.escape(module_list_attr)}\.(\d+)(\.|$)"
)
self._cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
self._cpu_dtypes: Dict[int, Dict[str, torch.dtype]] = {}
self._gpu_layers: Dict[int, Set[str]] = {}
self._named_parameters: Dict[str, torch.nn.Parameter] = {}
self._named_buffers: Dict[str, torch.Tensor] = {}
self._meta: Dict[str, Tuple[int, torch.dtype]] = {}
if auto_initialize:
self.initialize()
def _match_layer_idx(self, name: str) -> Optional[int]:
m = self._layer_name_re.search(name)
if not m:
return None
try:
return int(m.group(2))
except Exception:
return None
def _record_meta(self, name: str, t: torch.Tensor) -> None:
if name not in self._meta:
self._meta[name] = (int(t.ndim), t.dtype)
def _make_placeholder(self, name: str) -> torch.Tensor:
"""Rank-preserving empty placeholder on GPU."""
assert self.device is not None
ndim, dtype = self._meta[name]
shape = (0,) if ndim <= 0 else (0,) * ndim
return torch.empty(shape, device=self.device, dtype=dtype)
def _get_target(self, name: str) -> torch.Tensor:
if name in self._named_parameters:
return self._named_parameters[name]
return self._named_buffers[name]
def _offload_tensor(self, name: str, tensor: torch.Tensor, layer_idx: int) -> None:
if layer_idx not in self._cpu_weights:
self._cpu_weights[layer_idx] = {}
self._cpu_dtypes[layer_idx] = {}
self._record_meta(name, tensor)
cpu_weight = tensor.detach().to("cpu")
if self.pin_cpu_memory:
cpu_weight = cpu_weight.pin_memory()
self._cpu_weights[layer_idx][name] = cpu_weight
self._cpu_dtypes[layer_idx][name] = tensor.dtype
if self.device is not None:
tensor.data = self._make_placeholder(name)
@torch.compiler.disable
def initialize(self) -> None:
"""Offload all matched layer tensors to CPU and prefetch layer 0 (sync)."""
if not self.enabled:
return
self._named_parameters = dict(self.model.named_parameters())
self._named_buffers = dict(self.model.named_buffers())
for name, param in self._named_parameters.items():
layer_idx = self._match_layer_idx(name)
if layer_idx is None or layer_idx >= self.num_layers:
continue
self._offload_tensor(name, param, layer_idx)
for name, buf in self._named_buffers.items():
layer_idx = self._match_layer_idx(name)
if layer_idx is None or layer_idx >= self.num_layers:
continue
self._offload_tensor(name, buf, layer_idx)
self.prefetch_layer(0, non_blocking=False)
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
@torch.compiler.disable
def prefetch_layer(self, layer_idx: int, non_blocking: bool = True) -> None:
"""Prefetch a layer's tensors from CPU to GPU (async on copy_stream)."""
if not self.enabled or self.device is None or self.copy_stream is None:
return
if layer_idx < 0 or layer_idx >= self.num_layers:
return
if layer_idx in self._gpu_layers:
return
if layer_idx not in self._cpu_weights:
return
self.copy_stream.wait_stream(torch.cuda.current_stream())
param_names: Set[str] = set()
with torch.cuda.stream(self.copy_stream):
for name, cpu_weight in self._cpu_weights[layer_idx].items():
target = self._get_target(name)
gpu_weight = torch.empty(
cpu_weight.shape,
dtype=self._cpu_dtypes[layer_idx][name],
device=self.device,
)
gpu_weight.copy_(cpu_weight, non_blocking=non_blocking)
target.data = gpu_weight
param_names.add(name)
self._gpu_layers[layer_idx] = param_names
@contextmanager
def layer_scope(
self,
*,
prefetch_layer_idx: Optional[int],
release_layer_idx: Optional[int],
non_blocking: bool = True,
):
if self.enabled and release_layer_idx is not None:
cur = release_layer_idx
if (
cur not in self._gpu_layers
and cur in self._cpu_weights
and self.device is not None
and self.copy_stream is not None
):
self.prefetch_layer(cur, non_blocking=False)
torch.cuda.current_stream().wait_stream(self.copy_stream)
if self.enabled and prefetch_layer_idx is not None:
self.prefetch_layer(prefetch_layer_idx, non_blocking=non_blocking)
try:
yield
finally:
if self.enabled and self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
if self.enabled and release_layer_idx is not None:
self.release_layer(release_layer_idx)
@torch.compiler.disable
def release_layer(self, layer_idx: int) -> None:
"""Release a layer's tensors back to placeholders (free VRAM)."""
if not self.enabled or self.device is None:
return
if layer_idx < 0:
return
param_names = self._gpu_layers.pop(layer_idx, None)
if not param_names:
return
for name in param_names:
target = self._get_target(name)
# Ensure meta exists even if something unexpected happened
self._record_meta(name, target)
target.data = self._make_placeholder(name)
@torch.compiler.disable
def release_all(self) -> None:
"""Release all currently-resident layers back to placeholders."""
if not self.enabled or self.device is None:
return
if self.copy_stream is not None:
torch.cuda.current_stream().wait_stream(self.copy_stream)
for layer_idx in list(self._gpu_layers.keys()):
param_names = self._gpu_layers.pop(layer_idx, None)
if not param_names:
continue
for name in param_names:
target = self._get_target(name)
self._record_meta(name, target)
target.data = self._make_placeholder(name)
@@ -95,10 +95,6 @@ class ComponentLoader(ABC):
"image_encoder": (ImageEncoderLoader, "transformers"),
"upsampler": (UpsamplerLoader, "diffusers"),
"upsampler_2": (UpsamplerLoader, "diffusers"),
# Stable Audio's `StableAudioMultiConditioner` bundles T5 +
# NumberConditioners; not a pure text encoder, so it gets
# its own loader.
"conditioner": (ConditionerLoader, "fastvideo"),
}
if module_type in module_loaders:
@@ -1002,61 +998,6 @@ class SchedulerLoader(ComponentLoader):
return scheduler
class ConditionerLoader(ComponentLoader):
"""Loader for multi-conditioner components (e.g. Stable Audio's
`StableAudioMultiConditioner`, which bundles T5 + NumberConditioners
and is neither a pure text encoder nor a Diffusers-shaped module).
Reads `<subfolder>/config.json` to resolve the class via
`ModelRegistry`, instantiates with no args (the class pulls its own
defaults from its FastVideo config), then loads
`diffusion_pytorch_model.safetensors` non-strictly so externally
fetched sub-encoders (T5) don't trip the missing-key check.
"""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name", None)
config.pop("_name_or_path", None)
if class_name is None:
raise ValueError(
f"Conditioner config at {model_path} is missing the "
f"`_class_name` attribute required to resolve a model class.")
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
target_device = get_local_torch_device()
precision = getattr(fastvideo_args.pipeline_config, "precision", "fp16")
target_dtype = PRECISION_TO_TYPE.get(precision, torch.float16)
# Without this merge the model falls back to its dataclass
# defaults (e.g. SA-1.0's 3-conditioner spec — wrong for SA-small).
from dataclasses import fields as _fields
from fastvideo.configs.models.encoders import (
StableAudioConditionerConfig, )
if model_cls.__name__ == "StableAudioMultiConditioner":
cond_config = StableAudioConditionerConfig()
# `update_model_arch` is strict (raises on unknown keys); the
# converter writes a few non-arch keys (`_class_name`,
# `_diffusers_version`, `_name_or_path`) that must be filtered
# out first.
valid = {f.name for f in _fields(cond_config.arch_config)}
cond_config.update_model_arch({k: v for k, v in config.items() if k in valid})
with set_default_torch_dtype(target_dtype):
model = model_cls(cond_config)
else:
with set_default_torch_dtype(target_dtype):
model = model_cls()
weights = os.path.join(str(model_path), "diffusion_pytorch_model.safetensors")
if not os.path.isfile(weights):
raise FileNotFoundError(
f"Conditioner weights not found: {weights}")
state = safetensors_load_file(weights)
# Non-strict: T5 weights live outside this checkpoint (fetched in
# the conditioner's `__init__` from the standard HF repo).
model.load_state_dict(state, strict=False)
return model.to(device=target_device, dtype=target_dtype).eval()
class UpsamplerLoader(ComponentLoader):
"""Loader for upsamplers."""
-3
View File
@@ -86,9 +86,6 @@ _VAE_MODELS = {
("vaes", "gen3c_tokenizer_vae", "AutoencoderKLGen3CTokenizer"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
# `stable-audio-open-1.0/vae/config.json` ships `_class_name="AutoencoderOobleck"`
# (Diffusers' name); FastVideo's class is `OobleckVAE`.
"AutoencoderOobleck": ("vaes", "oobleck", "OobleckVAE"),
}
_AUDIO_MODELS = {
-376
View File
@@ -1,376 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 "Oobleck" VAE.
5-stage Conv1d autoencoder with Snake activations + diagonal-Gaussian
bottleneck. Loads `stabilityai/stable-audio-open-1.0/vae/` weights
directly via `OobleckVAE.from_pretrained(...)`.
vae = OobleckVAE.from_pretrained("stabilityai/stable-audio-open-1.0", subfolder="vae")
waveform = vae.decode(latent) # (B, audio_channels, samples)
latent = vae.encode(waveform).sample() # or .mode()
"""
from __future__ import annotations
import json
import math
import os
from dataclasses import dataclass
import numpy as np
import torch
import torch.nn as nn
from torch.nn.utils import weight_norm
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class Snake1d(nn.Module):
"""A 1D Snake activation with learnable per-channel alpha/beta."""
def __init__(self, hidden_dim: int, logscale: bool = True):
super().__init__()
self.alpha = nn.Parameter(torch.zeros(1, hidden_dim, 1))
self.beta = nn.Parameter(torch.zeros(1, hidden_dim, 1))
self.alpha.requires_grad = True
self.beta.requires_grad = True
self.logscale = logscale
def forward(self, x: torch.Tensor) -> torch.Tensor:
shape = x.shape
alpha = self.alpha if not self.logscale else torch.exp(self.alpha)
beta = self.beta if not self.logscale else torch.exp(self.beta)
x = x.reshape(shape[0], shape[1], -1)
x = x + (beta + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
return x.reshape(shape)
class OobleckResidualUnit(nn.Module):
def __init__(self, dimension: int = 16, dilation: int = 1):
super().__init__()
pad = ((7 - 1) * dilation) // 2
self.snake1 = Snake1d(dimension)
self.conv1 = weight_norm(nn.Conv1d(
dimension, dimension, kernel_size=7, dilation=dilation, padding=pad,
))
self.snake2 = Snake1d(dimension)
self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1))
def forward(self, x: torch.Tensor) -> torch.Tensor:
out = self.conv1(self.snake1(x))
out = self.conv2(self.snake2(out))
pad = (x.shape[-1] - out.shape[-1]) // 2
if pad > 0:
x = x[..., pad:-pad]
return x + out
class OobleckEncoderBlock(nn.Module):
def __init__(self, input_dim: int, output_dim: int, stride: int = 1):
super().__init__()
self.res_unit1 = OobleckResidualUnit(input_dim, dilation=1)
self.res_unit2 = OobleckResidualUnit(input_dim, dilation=3)
self.res_unit3 = OobleckResidualUnit(input_dim, dilation=9)
self.snake1 = Snake1d(input_dim)
self.conv1 = weight_norm(nn.Conv1d(
input_dim, output_dim,
kernel_size=2 * stride, stride=stride,
padding=math.ceil(stride / 2),
))
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.res_unit1(x)
x = self.res_unit2(x)
x = self.snake1(self.res_unit3(x))
return self.conv1(x)
class OobleckDecoderBlock(nn.Module):
def __init__(self, input_dim: int, output_dim: int, stride: int = 1):
super().__init__()
self.snake1 = Snake1d(input_dim)
self.conv_t1 = weight_norm(nn.ConvTranspose1d(
input_dim, output_dim,
kernel_size=2 * stride, stride=stride,
padding=math.ceil(stride / 2),
))
self.res_unit1 = OobleckResidualUnit(output_dim, dilation=1)
self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3)
self.res_unit3 = OobleckResidualUnit(output_dim, dilation=9)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.snake1(x)
x = self.conv_t1(x)
x = self.res_unit1(x)
x = self.res_unit2(x)
return self.res_unit3(x)
class OobleckDiagonalGaussianDistribution:
"""Diagonal-Gaussian VAE posterior with `softplus(scale) + 1e-4` std."""
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
self.parameters = parameters
self.mean, self.scale = parameters.chunk(2, dim=1)
self.std = nn.functional.softplus(self.scale) + 1e-4
self.var = self.std * self.std
self.logvar = torch.log(self.var)
self.deterministic = deterministic
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
noise = torch.randn(
self.mean.shape, generator=generator,
device=self.parameters.device, dtype=self.parameters.dtype,
)
return self.mean + self.std * noise
def mode(self) -> torch.Tensor:
return self.mean
@dataclass
class OobleckDecoderOutput:
sample: torch.Tensor
class OobleckEncoder(nn.Module):
def __init__(
self,
encoder_hidden_size: int,
audio_channels: int,
downsampling_ratios: list[int],
channel_multiples: list[int],
):
super().__init__()
strides = downsampling_ratios
channel_multiples = [1] + list(channel_multiples)
self.conv1 = weight_norm(nn.Conv1d(
audio_channels, encoder_hidden_size, kernel_size=7, padding=3,
))
self.block = nn.ModuleList([
OobleckEncoderBlock(
input_dim=encoder_hidden_size * channel_multiples[i],
output_dim=encoder_hidden_size * channel_multiples[i + 1],
stride=s,
)
for i, s in enumerate(strides)
])
d_model = encoder_hidden_size * channel_multiples[-1]
self.snake1 = Snake1d(d_model)
self.conv2 = weight_norm(nn.Conv1d(
d_model, encoder_hidden_size, kernel_size=3, padding=1,
))
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.conv1(x)
for m in self.block:
x = m(x)
x = self.snake1(x)
return self.conv2(x)
class OobleckDecoder(nn.Module):
def __init__(
self,
channels: int,
input_channels: int,
audio_channels: int,
upsampling_ratios: list[int],
channel_multiples: list[int],
):
super().__init__()
strides = upsampling_ratios
channel_multiples = [1] + list(channel_multiples)
self.conv1 = weight_norm(nn.Conv1d(
input_channels, channels * channel_multiples[-1],
kernel_size=7, padding=3,
))
self.block = nn.ModuleList([
OobleckDecoderBlock(
input_dim=channels * channel_multiples[len(strides) - i],
output_dim=channels * channel_multiples[len(strides) - i - 1],
stride=s,
)
for i, s in enumerate(strides)
])
self.snake1 = Snake1d(channels)
self.conv2 = weight_norm(nn.Conv1d(
channels, audio_channels, kernel_size=7, padding=3, bias=False,
))
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.conv1(x)
for layer in self.block:
x = layer(x)
x = self.snake1(x)
return self.conv2(x)
# ---------------------------------------------------------------------------
# Top-level VAE
# ---------------------------------------------------------------------------
class OobleckVAE(nn.Module):
"""Stable Audio Open 1.0 VAE.
Constructed either from an `OobleckVAEConfig` (the standard
`VAELoader` path) or from explicit kwargs (back-compat for tests
and `from_pretrained` callers).
"""
def __init__(
self,
config=None, # type: OobleckVAEConfig | None
*,
encoder_hidden_size: int = 128,
downsampling_ratios: list[int] | None = None,
channel_multiples: list[int] | None = None,
decoder_channels: int = 128,
decoder_input_channels: int = 64,
audio_channels: int = 2,
sampling_rate: int = 44100,
):
super().__init__()
if config is not None:
arch = config.arch_config
encoder_hidden_size = arch.encoder_hidden_size
downsampling_ratios = list(arch.downsampling_ratios)
channel_multiples = list(arch.channel_multiples)
decoder_channels = arch.decoder_channels
decoder_input_channels = arch.decoder_input_channels
audio_channels = arch.audio_channels
sampling_rate = arch.sampling_rate
if downsampling_ratios is None:
downsampling_ratios = [2, 4, 4, 8, 8]
if channel_multiples is None:
channel_multiples = [1, 2, 4, 8, 16]
self.encoder_hidden_size = encoder_hidden_size
self.downsampling_ratios = downsampling_ratios
self.decoder_channels = decoder_channels
self.upsampling_ratios = list(reversed(downsampling_ratios))
self.hop_length = int(np.prod(downsampling_ratios))
self.sampling_rate = sampling_rate
self.audio_channels = audio_channels
self.decoder_input_channels = decoder_input_channels
self.encoder = OobleckEncoder(
encoder_hidden_size=encoder_hidden_size,
audio_channels=audio_channels,
downsampling_ratios=downsampling_ratios,
channel_multiples=channel_multiples,
)
self.decoder = OobleckDecoder(
channels=decoder_channels,
input_channels=decoder_input_channels,
audio_channels=audio_channels,
upsampling_ratios=self.upsampling_ratios,
channel_multiples=channel_multiples,
)
def encode(
self, x: torch.Tensor,
) -> OobleckDiagonalGaussianDistribution:
return OobleckDiagonalGaussianDistribution(self.encoder(x))
def decode(self, z: torch.Tensor) -> OobleckDecoderOutput:
return OobleckDecoderOutput(sample=self.decoder(z))
def forward(
self, sample: torch.Tensor, sample_posterior: bool = False,
) -> OobleckDecoderOutput:
posterior = self.encode(sample)
z = posterior.sample() if sample_posterior else posterior.mode()
return self.decode(z)
# -------------------------------------------------------------------
# Loader
# -------------------------------------------------------------------
@classmethod
def from_pretrained(
cls,
model_path: str,
*,
subfolder: str | None = None,
torch_dtype: torch.dtype | None = None,
) -> "OobleckVAE":
"""Instantiate and load weights from a Stable Audio VAE dir.
`model_path` may be:
* a HF repo id (e.g. `stabilityai/stable-audio-open-1.0`),
* a local directory containing `config.json` + safetensors,
* a local directory whose `subfolder="vae"` holds those files.
For gated repos, the HF token is read from `HF_TOKEN` /
`HUGGINGFACE_HUB_TOKEN` / `HF_API_KEY` (see `resolve_hf_token`).
"""
import inspect
from safetensors.torch import load_file
from fastvideo.utils import resolve_hf_token
# Resolve to a local directory.
if os.path.isdir(model_path):
root = model_path
else:
from huggingface_hub import snapshot_download
allow = ["vae/*"] if subfolder else ["*"]
root = snapshot_download(
repo_id=model_path, token=resolve_hf_token(), allow_patterns=allow,
)
if subfolder:
root = os.path.join(root, subfolder)
if not os.path.isdir(root):
raise FileNotFoundError(f"Not a directory: {root}")
cfg_path = os.path.join(root, "config.json")
if not os.path.isfile(cfg_path):
raise FileNotFoundError(
f"Expected config.json at {cfg_path}. If using a HF repo, "
f"pass subfolder='vae'."
)
with open(cfg_path) as f:
cfg = json.load(f)
cfg_fields = {k: v for k, v in cfg.items() if not k.startswith("_")}
# Diffusers configs commonly carry extra fields (`scaling_factor`,
# `_diffusers_version`, ...) the bare `OobleckVAE` ctor doesn't accept.
init_params = inspect.signature(cls.__init__).parameters
cfg_fields = {k: v for k, v in cfg_fields.items() if k in init_params}
model = cls(**cfg_fields)
weights_path = os.path.join(root, "diffusion_pytorch_model.safetensors")
if not os.path.isfile(weights_path):
# Allow `model.safetensors` as a fallback.
alt = os.path.join(root, "model.safetensors")
if os.path.isfile(alt):
weights_path = alt
else:
raise FileNotFoundError(
f"No safetensors weights under {root}. Expected "
f"diffusion_pytorch_model.safetensors."
)
state = load_file(weights_path)
missing, unexpected = model.load_state_dict(state, strict=False)
if missing:
raise RuntimeError(
f"OobleckVAE missing {len(missing)} keys from {weights_path}: "
f"{missing[:5]}"
)
if unexpected:
# Non-critical: some checkpoints embed the VAE inside a larger
# container (e.g. `pretransform.model.*`). Log the count so
# genuine loader regressions don't go unnoticed.
logger.debug(
"OobleckVAE: ignored %d unexpected keys from %s "
"(first 3: %s)", len(unexpected), weights_path, unexpected[:3],
)
if torch_dtype is not None:
model = model.to(dtype=torch_dtype)
model.eval()
return model
EntryClass = OobleckVAE
-122
View File
@@ -1,122 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Lazy-loading pipeline wrapper around `OobleckVAE`.
Two reasons this exists rather than using `OobleckVAE` directly:
1. The underlying VAE is fetched on first `encode`/`decode` call, not
at construction — lets pipelines build the module tree on CPU
before knowing the target device.
2. The lazy VAE's params are hidden from `named_parameters()` so the
FastVideo pipeline-component loader doesn't try to match Oobleck's
safetensors against the host pipeline's converted-repo state dict.
For standalone use prefer `OobleckVAE.from_pretrained(...)` directly.
"""
from __future__ import annotations
import os
import torch
from torch import nn
from fastvideo.configs.models.vaes import OobleckVAEConfig
class SAAudioVAEModel(nn.Module):
"""Pipeline-glue lazy loader around `OobleckVAE`."""
def __init__(self, config: OobleckVAEConfig) -> None:
super().__init__()
self.config = config
arch = config.arch_config
self.pretrained_path: str = config.pretrained_path
self.pretrained_subfolder: str | None = config.pretrained_subfolder
self.pretrained_dtype: str = config.pretrained_dtype
self.sampling_rate: int = arch.sampling_rate
self.audio_channels: int = arch.audio_channels
self.decoder_input_channels: int = arch.decoder_input_channels
self._oobleck_vae = None
def named_parameters(self, prefix: str = "", recurse: bool = True):
# Hide the lazy-loaded VAE — its weights are fetched separately
# and shouldn't appear in the host pipeline's loader sweep.
for name, param in super().named_parameters(prefix=prefix, recurse=recurse):
if name.startswith("_oobleck_vae.") or name == "_oobleck_vae":
continue
yield name, param
def _build(self, device: torch.device | None = None):
from fastvideo.models.vaes.oobleck import OobleckVAE
path = self.pretrained_path
if not path:
raise ValueError(
"OobleckVAEConfig.pretrained_path must be set; expected "
"`stabilityai/stable-audio-open-1.0` or a local path."
)
dtype = getattr(torch, self.pretrained_dtype, torch.float32)
# If the caller already pointed us at the VAE dir directly, drop
# the subfolder. Otherwise pass through (default "vae").
subfolder: str | None = self.pretrained_subfolder
if subfolder and os.path.isdir(path) and os.path.isfile(os.path.join(path, "config.json")):
subfolder = None
model = OobleckVAE.from_pretrained(path, subfolder=subfolder, torch_dtype=dtype)
if device is not None:
model = model.to(device=device)
model.eval()
return model
@property
def oobleck_vae(self):
if self._oobleck_vae is None:
self._oobleck_vae = self._build()
return self._oobleck_vae
# Back-compat alias: callers that imported this earlier referred to
# the underlying VAE as `sa_audio_vae_model`. Both names point at the
# same object.
@property
def sa_audio_vae_model(self):
return self.oobleck_vae
@property
def hop_length(self) -> int:
return int(self.oobleck_vae.hop_length)
def _move_to_input_device(self, model, ref: torch.Tensor):
if ref is None:
return model
first_param = next(model.parameters(), None)
if first_param is not None and first_param.device != ref.device:
model = model.to(device=ref.device)
self._oobleck_vae = model
return model
def decode(self, latent: torch.Tensor) -> torch.Tensor:
"""Decode an audio latent (`[B, C_latent, L]`) -> waveform
(`[B, audio_channels, samples]`).
"""
model = self.oobleck_vae
model = self._move_to_input_device(model, latent)
with torch.no_grad():
out = model.decode(latent.to(next(model.parameters()).dtype))
if hasattr(out, "sample"):
return out.sample
return out
def encode(self, waveform: torch.Tensor, sample_posterior: bool = False) -> torch.Tensor:
"""Encode `[B, C_audio, samples]` -> latent `[B, C_latent, L]`.
`sample_posterior=False` (default): deterministic mean.
`sample_posterior=True`: stochastic sample (`mean + softplus(scale) * randn`).
"""
model = self.oobleck_vae
model = self._move_to_input_device(model, waveform)
with torch.no_grad():
out = model.encode(waveform.to(next(model.parameters()).dtype))
if hasattr(out, "latent_dist"):
out = out.latent_dist
return out.sample() if sample_posterior else out.mode()
EntryClass = SAAudioVAEModel
@@ -22,18 +22,23 @@ class Cosmos2VideoToWorldPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler", "safety_checker"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
scheduler = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
use_karras_sigmas=True,
)
scheduler.config.sigma_max = 80.0
scheduler.config.sigma_min = 0.002
scheduler.config.sigma_data = 1.0
scheduler.config.final_sigmas_type = "sigma_min"
scheduler.sigma_max = 80.0
scheduler.sigma_min = 0.002
scheduler.sigma_data = 1.0
self.modules["scheduler"] = scheduler
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift,
use_karras_sigmas=True)
sigma_max = 80.0
sigma_min = 0.002
sigma_data = 1.0
final_sigmas_type = "sigma_min"
if self.modules["scheduler"] is not None:
scheduler = self.modules["scheduler"]
scheduler.config.sigma_max = sigma_max
scheduler.config.sigma_min = sigma_min
scheduler.config.sigma_data = sigma_data
scheduler.config.final_sigmas_type = final_sigmas_type
scheduler.sigma_max = sigma_max
scheduler.sigma_min = sigma_min
scheduler.sigma_data = sigma_data
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
@@ -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,
},
)
@@ -1 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
@@ -1,63 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio presets.
Sampling defaults track the published HF model card
(https://huggingface.co/stabilityai/stable-audio-open-1.0):
100 steps, CFG=7, dpmpp-3m-sde, sigma_min=0.3, sigma_max=500, rho=1.0.
"""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="Stable Audio Cosine-DPM++ denoising with text + duration CFG.",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
# `audio_start_in_s` / `audio_end_in_s` are call-kwargs, kept off here.
# `height`/`width` are pinned to 8 (the shared `InputValidationStage`
# rejects values that aren't divisible by 8) and `num_frames` to 1 so
# the video-shaped preallocation in `VideoGenerator` stays tiny — the
# real output is the audio waveform on `result["audio"]`, not the
# placeholder frame tensor.
_SHARED_DEFAULTS = {
"seed": 0,
"guidance_scale": 7.0,
"num_inference_steps": 100,
"negative_prompt": "",
"height": 8,
"width": 8,
"num_frames": 1,
}
STABLE_AUDIO_OPEN_1_0_BASE = InferencePreset(
name="stable_audio_open_1_0_base",
version=1,
model_family="stable_audio",
description=("Stability AI Stable Audio Open 1.0 text-to-audio. Generates up "
"to ~47.5s of stereo 44.1 kHz audio per call. Default duration "
"is 10s; raise via `audio_end_in_s` up to the model max."),
workload_type="t2v", # NOTE: WorkloadType has no T2A variant yet (REVIEW item 28)
stage_schemas=(_DENOISE_STAGE, ),
defaults=dict(_SHARED_DEFAULTS),
)
# Smaller / faster checkpoint with the same Oobleck VAE but a 1024-dim
# 16-layer DiT with `qk_norm="ln"`. Sampling defaults match the official
# `stable-audio-open-small` model card.
STABLE_AUDIO_OPEN_SMALL = InferencePreset(
name="stable_audio_open_small",
version=1,
model_family="stable_audio",
description=("Stability AI Stable Audio Open Small text-to-audio. Faster than "
"the 1.0 base; supports up to ~11.9s of stereo 44.1 kHz audio per "
"call (smaller training window)."),
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults=dict(_SHARED_DEFAULTS),
)
ALL_PRESETS = (STABLE_AUDIO_OPEN_1_0_BASE, STABLE_AUDIO_OPEN_SMALL)
@@ -1,125 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio Open 1.0 pipeline (T2A + A2A + RePaint inpainting).
Components are loaded via the standard
`ComposedPipelineBase.load_modules` against the FastVideo-curated
Diffusers-format repo `FastVideo/stable-audio-open-1.0-Diffusers`
(produced by
`scripts/checkpoint_conversion/stable_audio_to_diffusers.py`). The DiT
is a `BaseDiT` subclass loaded by `TransformerLoader`; the VAE is
loaded by `VAELoader`; the multi-conditioner (T5 + NumberConditioners)
is loaded by `ConditionerLoader` (a Stable Audio-specific addition).
Stages:
InputValidationStage
→ StableAudioConditioningStage (T5 + NumberConditioner -> cross-attn + global cond, with CFG)
→ StableAudioLatentPreparationStage (initial Gaussian noise; encodes A2A / inpaint refs)
→ StableAudioDenoisingStage (k-diffusion `dpmpp-3m-sde` over the DiT)
→ StableAudioDecodingStage (OobleckVAE -> waveform)
"""
from __future__ import annotations
import functools
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.basic.stable_audio.stages import (
StableAudioConditioningStage,
StableAudioDecodingStage,
StableAudioDenoisingStage,
StableAudioLatentPreparationStage,
)
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import InputValidationStage
logger = init_logger(__name__)
@functools.lru_cache(maxsize=1)
def _warn_tf32_disabled_for_stable_audio() -> None:
logger.warning("Stable Audio pipeline is disabling process-global "
"torch.backends.{cuda.matmul.allow_tf32, cudnn.allow_tf32, "
"cuda.matmul.allow_fp16_reduced_precision_reduction, "
"cudnn.benchmark} for A2A renoise determinism. Other models "
"loaded into this process will inherit these settings.")
def _disable_tf32_for_stable_audio() -> None:
"""Disable TF32 / cuDNN nondeterminism — A2A renoise-then-denoise SDE
amplifies per-element drift, and the published parity bounds were
set with these off. Process-global; the first call logs a warning.
"""
_warn_tf32_disabled_for_stable_audio()
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction = False
torch.backends.cudnn.benchmark = False
class StableAudioPipeline(ComposedPipelineBase):
"""Stable Audio Open 1.0 pipeline.
Mode is kwargs-driven on `generate_video()`:
* Text-to-audio (default) -- `prompt=...`, `audio_end_in_s=...`
* Audio-to-audio variation -- add `init_audio=ref` (and optionally
`init_noise_level`, lower = closer to reference)
* RePaint inpainting / outpainting -- add `inpaint_audio=ref` and
`inpaint_mask` (1-D, 1 = keep / 0 = regenerate)
See `examples/inference/basic/basic_stable_audio*.py` for runnable
examples of each mode.
"""
_required_config_modules = [
"vae",
"transformer",
"conditioner",
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
"""Apply Stable Audio's process-global numerics overrides BEFORE
the standard component loaders run (TF32 off for A2A renoise
determinism)."""
_disable_tf32_for_stable_audio()
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
pc = fastvideo_args.pipeline_config
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(
stage_name="conditioning_stage",
stage=StableAudioConditioningStage(conditioner=self.get_module("conditioner")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=StableAudioLatentPreparationStage(
io_channels=64,
# Per-variant training window: 2,097,152 (~47.5s) for
# SA-1.0; 524,288 (~11.9s) for SA-small. Pulled from the
# pipeline config so each variant gets its own latent
# length.
sample_size=pc.sample_size,
vae=self.get_module("vae"),
sample_rate=pc.sampling_rate,
audio_channels=pc.audio_channels,
),
)
self.add_stage(
stage_name="denoising_stage",
stage=StableAudioDenoisingStage(transformer=self.get_module("transformer")),
)
self.add_stage(
stage_name="decoding_stage",
stage=StableAudioDecodingStage(vae=self.get_module("vae")),
)
EntryClass = StableAudioPipeline
@@ -1,12 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.pipelines.basic.stable_audio.stages.conditioning import StableAudioConditioningStage
from fastvideo.pipelines.basic.stable_audio.stages.decoding import StableAudioDecodingStage
from fastvideo.pipelines.basic.stable_audio.stages.denoising import StableAudioDenoisingStage
from fastvideo.pipelines.basic.stable_audio.stages.latent_preparation import StableAudioLatentPreparationStage
__all__ = [
"StableAudioConditioningStage",
"StableAudioDecodingStage",
"StableAudioDenoisingStage",
"StableAudioLatentPreparationStage",
]
@@ -1,98 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio conditioning stage."""
from __future__ import annotations
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
class StableAudioConditioningStage(PipelineStage):
"""Run the conditioner over the prompt + duration and stash the
DiT-ready (cross_attn_cond, cross_attn_mask, global_embed) triple
on `batch.extra` (plus the negative-prompt triple when CFG is on).
"""
def __init__(self, conditioner) -> None:
super().__init__()
self.conditioner = conditioner
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
pc = fastvideo_args.pipeline_config
device = next(self.conditioner.parameters()).device
start_attr = getattr(batch, "audio_start_in_s", None)
end_attr = getattr(batch, "audio_end_in_s", None)
audio_start_in_s = float(start_attr if start_attr is not None else pc.audio_start_in_s)
audio_end_in_s = float(end_attr if end_attr is not None else pc.audio_end_in_s)
max_duration = float(getattr(pc, "max_audio_duration_s", 2097152 / 44100))
if audio_start_in_s < 0:
raise ValueError(f"audio_start_in_s must be >= 0, got {audio_start_in_s}.")
if audio_end_in_s <= audio_start_in_s:
raise ValueError(f"audio_end_in_s ({audio_end_in_s}) must be > audio_start_in_s "
f"({audio_start_in_s}).")
if audio_end_in_s > max_duration:
raise ValueError(f"audio_end_in_s ({audio_end_in_s}s) exceeds the model's fixed "
f"window of {max_duration:.4f}s. Stable Audio Open 1.0 always "
f"samples a 2,097,152-frame latent and slices to "
f"[start, end] after decode; values past the window are silently "
f"truncated. Lower audio_end_in_s or split the request.")
guidance_scale = float(batch.guidance_scale or pc.guidance_scale)
do_cfg = guidance_scale > 1.0
if isinstance(batch.prompt, str):
prompt = batch.prompt
elif isinstance(batch.prompt, list):
if len(batch.prompt) > 1:
raise ValueError(f"Stable Audio does not support batched prompts; got "
f"{len(batch.prompt)} entries. Pass a single string or a "
f"single-element list.")
prompt = batch.prompt[0] if batch.prompt else ""
else:
raise TypeError(f"`prompt` must be a string or a list of strings, got "
f"{type(batch.prompt).__name__}.")
# Send only the keys the conditioner declares (per-variant).
all_cond_values = {
"prompt": prompt,
"seconds_start": audio_start_in_s,
"seconds_total": audio_end_in_s,
}
active_ids = self.conditioner.cross_attention_cond_ids
cond_meta = [{k: all_cond_values[k] for k in active_ids if k in all_cond_values}]
cond = self.conditioner(cond_meta, device)
cross_attn_cond, cross_attn_mask, global_embed = self.conditioner.get_conditioning_inputs(cond)
neg_cross_attn_cond = None
neg_cross_attn_mask = None
neg_global_embed = None
if do_cfg:
neg_prompt = batch.negative_prompt or ""
if isinstance(neg_prompt, list):
neg_prompt = neg_prompt[0] if neg_prompt else ""
neg_values = dict(all_cond_values, prompt=neg_prompt)
neg_meta = [{k: neg_values[k] for k in active_ids if k in neg_values}]
neg = self.conditioner(neg_meta, device)
neg_cross_attn_cond, neg_cross_attn_mask, neg_global_embed = (self.conditioner.get_conditioning_inputs(neg))
if batch.extra is None:
batch.extra = {}
batch.extra["cross_attn_cond"] = cross_attn_cond
batch.extra["cross_attn_mask"] = cross_attn_mask
batch.extra["global_embed"] = global_embed
batch.extra["negative_cross_attn_cond"] = neg_cross_attn_cond
batch.extra["negative_cross_attn_mask"] = neg_cross_attn_mask
batch.extra["negative_global_embed"] = neg_global_embed
batch.extra["do_cfg"] = do_cfg
batch.extra["audio_start_in_s"] = audio_start_in_s
batch.extra["audio_end_in_s"] = audio_end_in_s
return batch
@@ -1,67 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio decoding: latent -> waveform via OobleckVAE.
Slices the output to `[audio_start_in_s, audio_end_in_s]` and stashes
the result on `batch.extra["audio"]` + `["audio_sample_rate"]` for
`VideoGenerator._mux_audio` to pick up.
"""
from __future__ import annotations
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
class StableAudioDecodingStage(PipelineStage):
"""Decode latent → audio waveform + slice to [start, end]."""
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
pc = fastvideo_args.pipeline_config
latents = batch.latents
# VAE may be CPU-parked under `vae_cpu_offload=True`.
from fastvideo.distributed.parallel_state import get_local_torch_device
self.vae = self.vae.to(get_local_torch_device())
decoded = self.vae.decode(latents)
if hasattr(decoded, "sample"): # tolerate tensor or dataclass
decoded = decoded.sample
sr = int(getattr(self.vae, "sampling_rate", pc.sampling_rate))
start_in_s = float(batch.extra.get("audio_start_in_s", pc.audio_start_in_s))
end_in_s = float(batch.extra.get("audio_end_in_s", pc.audio_end_in_s))
decoded = decoded[:, :, int(start_in_s * sr):int(end_in_s * sr)]
if batch.extra is None:
batch.extra = {}
# `_mux_audio` / `_write_pcm_wav` want `[samples, channels]`.
batch.extra["audio"] = decoded.squeeze(0).T.detach().float().cpu().numpy()
batch.extra["audio_sample_rate"] = sr
batch.extra["audio_only"] = True
# Raw tensor for parity tests.
batch.extra["decoded_audio"] = decoded.detach().cpu()
# `VideoGenerator.generate_video` is video-shaped (asserts
# `output_batch.output is not None`); fill with a placeholder of
# the expected `[B, 3, num_frames, H, W]` shape — the real audio
# is on `batch.extra` above. Pure-audio workload support tracked
# in REVIEW item 28.
b = decoded.shape[0]
n_frames = int(getattr(batch, "num_frames", 1) or 1)
h = int(getattr(batch, "height", 1) or 1)
w = int(getattr(batch, "width", 1) or 1)
batch.output = torch.zeros((b, 3, n_frames, h, w), dtype=torch.uint8)
return batch
@@ -1,207 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio denoising — k-diffusion `dpmpp-3m-sde` over the DiT.
CFG-batched conditioning is built once outside the sampler loop so the
adapter only does `cat([x, x])` + DiT call per step.
"""
from __future__ import annotations
import math
import torch
import torch.nn as nn
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
class _DiTAdapter(nn.Module):
"""`StableAudioDiT` -> `K.external.VDenoiser` adapter.
`batch_cond` / `batch_global` are precomputed CFG-batched tensors
(`[2, ...]` for CFG, `[1, ...]` otherwise); building them once
outside the sampler loop saves ~3 cats × 100 steps per call.
"""
def __init__(self, dit, *, batch_cond: torch.Tensor, batch_global: torch.Tensor, cfg_scale: float) -> None:
super().__init__()
self.dit = dit
self.batch_cond = batch_cond
self.batch_global = batch_global
self.cfg_scale = cfg_scale
self.do_cfg = cfg_scale != 1.0
def forward(self, x: torch.Tensor, t: torch.Tensor, **_unused) -> torch.Tensor:
if not self.do_cfg:
return self.dit(x, t, cross_attn_cond=self.batch_cond, global_embed=self.batch_global)
batch_x = torch.cat([x, x], dim=0)
batch_t = torch.cat([t, t], dim=0)
out = self.dit(batch_x, batch_t, cross_attn_cond=self.batch_cond, global_embed=self.batch_global)
cond_out, uncond_out = torch.chunk(out, 2, dim=0)
return uncond_out + (cond_out - uncond_out) * self.cfg_scale
class StableAudioDenoisingStage(PipelineStage):
"""k-diffusion `dpmpp-3m-sde` sampling loop."""
# Sampler defaults from the published model card.
_SIGMA_MIN = 0.3
_SIGMA_MAX = 500.0
_RHO = 1.0
_LOG_SIGMA_MIN = math.log(_SIGMA_MIN)
_LOG_SIGMA_MAX = math.log(_SIGMA_MAX)
def __init__(self, transformer) -> None:
super().__init__()
self.transformer = transformer
def _resolve_sigma_max(self, batch) -> float:
"""Map A2A intent to `sigma_max`.
Public knob is `init_audio_strength` (0..1, higher = closer to
source), log-interpolated between SIGMA_MIN (= preservation) and
SIGMA_MAX (= full T2A). Raw `init_noise_level` is the legacy
sigma_max override; passing both is an error.
"""
raw = getattr(batch, "init_noise_level", None)
strength = getattr(batch, "init_audio_strength", None)
if raw is not None and strength is not None:
raise ValueError("Pass `init_audio_strength` (0..1) OR `init_noise_level` "
"(raw sigma_max), not both.")
if raw is not None:
return float(raw)
s = max(0.0, min(1.0, float(strength) if strength is not None else 0.6))
return float(math.exp(self._LOG_SIGMA_MAX - s * (self._LOG_SIGMA_MAX - self._LOG_SIGMA_MIN)))
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
pc = fastvideo_args.pipeline_config
ext = batch.extra
device = batch.latents.device
guidance_scale = float(batch.guidance_scale or pc.guidance_scale)
steps = int(batch.num_inference_steps)
import k_diffusion as K
init_latent = ext.get("init_latent")
sigma_max = self._resolve_sigma_max(batch) if init_latent is not None else self._SIGMA_MAX
sigmas = K.sampling.get_sigmas_polyexponential(steps, self._SIGMA_MIN, sigma_max, self._RHO, device=device)
# Cast noise + conditioning to the DiT's dtype before sampling
# (matches `stable_audio_tools/inference/generation.py:185-187`).
model_dtype = next(self.transformer.parameters()).dtype
def _cast(t: torch.Tensor | None) -> torch.Tensor | None:
return t.to(model_dtype) if t is not None else None
x = (batch.latents * sigmas[0]).to(model_dtype)
if init_latent is not None:
x = x + init_latent.to(model_dtype)
batch_cond, batch_global = _build_cfg_conditioning(
cross_attn_cond=ext["cross_attn_cond"].to(model_dtype),
global_embed=ext["global_embed"].to(model_dtype),
negative_cross_attn_cond=_cast(ext.get("negative_cross_attn_cond")),
negative_cross_attn_mask=ext.get("negative_cross_attn_mask"),
negative_global_embed=_cast(ext.get("negative_global_embed")),
do_cfg=guidance_scale != 1.0,
)
adapter = _DiTAdapter(self.transformer,
batch_cond=batch_cond,
batch_global=batch_global,
cfg_scale=guidance_scale)
denoiser = K.external.VDenoiser(adapter)
# RePaint blending hook — works on any v-prediction model, no
# inpaint-trained checkpoint needed.
inpaint_mask = ext.get("inpaint_mask_latent")
inpaint_ref = ext.get("inpaint_reference_latent")
if inpaint_mask is not None and inpaint_ref is not None:
inpaint_mask = inpaint_mask.to(model_dtype)
inpaint_ref = inpaint_ref.to(model_dtype)
callback = _make_inpaint_callback(inpaint_ref, inpaint_mask, sigmas)
else:
callback = None
# `LocalAttention` (in `StableAudioDiT`) reads `get_forward_context()`
# for `attn_metadata`; wrap the whole loop.
with set_forward_context(current_timestep=0, attn_metadata=None):
sampled = K.sampling.sample_dpmpp_3m_sde(denoiser,
x,
sigmas,
disable=False,
extra_args={},
callback=callback)
# Final blend so the kept region of the inpaint reference is exact.
if inpaint_mask is not None and inpaint_ref is not None:
sampled = inpaint_ref * inpaint_mask + sampled * (1 - inpaint_mask)
batch.latents = sampled
return batch
def _build_cfg_conditioning(
*,
cross_attn_cond: torch.Tensor,
global_embed: torch.Tensor,
negative_cross_attn_cond: torch.Tensor | None,
negative_cross_attn_mask: torch.Tensor | None,
negative_global_embed: torch.Tensor | None,
do_cfg: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Build the CFG-batched `(cond, global)` tensors once.
Cond ordering is `[conditioned, unconditioned]` (the adapter splits
with the same convention). Masked negative cond is zero-filled where
`mask == 0`.
"""
if not do_cfg:
return cross_attn_cond, global_embed
if negative_cross_attn_cond is not None:
if negative_cross_attn_mask is not None:
neg_mask = negative_cross_attn_mask.to(torch.bool).unsqueeze(2)
null_embed = torch.zeros_like(cross_attn_cond)
negative_cross_attn_cond = torch.where(neg_mask, negative_cross_attn_cond, null_embed)
batch_cond = torch.cat([cross_attn_cond, negative_cross_attn_cond], dim=0)
else:
batch_cond = torch.cat([cross_attn_cond, torch.zeros_like(cross_attn_cond)], dim=0)
other_global = global_embed if negative_global_embed is None else negative_global_embed
batch_global = torch.cat([global_embed, other_global], dim=0)
return batch_cond, batch_global
def _make_inpaint_callback(reference_latent: torch.Tensor, mask: torch.Tensor, sigmas: torch.Tensor):
"""RePaint blending callback for the k-diffusion sampler.
At every step, replaces the kept region (`mask == 1`) of the in-
flight latent with the reference re-noised to the next sigma —
pulls the kept region back onto the trajectory the model expects,
so RePaint-style inpainting converges on non-inpaint-trained models.
Pre-allocates the noise buffer so the ~100 sampler steps don't churn
~25 MB of fresh allocations per call.
"""
noise_buf = torch.empty_like(reference_latent)
inv_mask = 1 - mask
def cb(info: dict) -> None:
i = int(info["i"])
next_i = min(i + 1, len(sigmas) - 1)
sigma_next = float(sigmas[next_i])
noise_buf.normal_()
# `state["x"]` is the live latent; the dpmpp-3m-sde sampler picks
# up our in-place mutation between steps (verified against
# k_diffusion 0.1.1.post1).
x = info["x"]
x.copy_((reference_latent + noise_buf * sigma_next) * mask + x * inv_mask)
return cb
@@ -1,174 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Stable Audio latent preparation.
Seeds + samples the initial Gaussian noise; encodes `init_audio` (A2A
variation) or `inpaint_audio` + `inpaint_mask` (RePaint inpainting) into
latent-space tensors on `batch.extra` for the denoising stage.
"""
from __future__ import annotations
import os
import torch
import torch.nn.functional as F
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
class StableAudioLatentPreparationStage(PipelineStage):
def __init__(self,
io_channels: int = 64,
sample_size: int = 2097152,
vae=None,
sample_rate: int = 44100,
audio_channels: int = 2) -> None:
super().__init__()
self.io_channels = io_channels
# Audio-domain length the model was trained for; latent length
# = sample_size // vae.hop_length (= 2097152 / 2048 = 1024).
self.sample_size = sample_size
self.vae = vae # used to encode init_audio / inpaint_audio
self.sample_rate = sample_rate
self.audio_channels = audio_channels
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
ext = batch.extra or {}
device = ext["cross_attn_cond"].device
latent_sample_size = self.sample_size // self._hop_length()
seed = int(batch.seed) if batch.seed is not None else 0
torch.manual_seed(seed)
latents = torch.randn((1, self.io_channels, latent_sample_size), device=device)
batch.latents = latents
if batch.extra is None:
batch.extra = {}
init_audio = getattr(batch, "init_audio", None)
inpaint_audio = getattr(batch, "inpaint_audio", None)
inpaint_mask = getattr(batch, "inpaint_mask", None)
# Loud-fail rather than silently falling through to T2A.
if inpaint_audio is not None and inpaint_mask is None:
raise ValueError("Stable Audio inpainting requires both `inpaint_audio` and "
"`inpaint_mask` (1-D tensor in {0, 1} at the model sample rate, "
"1 = keep, 0 = regenerate). Got `inpaint_audio` without `inpaint_mask`.")
if inpaint_mask is not None and inpaint_audio is None:
raise ValueError("Stable Audio inpainting requires both `inpaint_audio` and "
"`inpaint_mask`. Got `inpaint_mask` without `inpaint_audio` — "
"did you mean to pass `init_audio` (audio-to-audio variation)?")
if init_audio is not None and inpaint_audio is not None:
raise ValueError("Stable Audio cannot do A2A variation and inpainting in the "
"same call. Pass either `init_audio` (variation) or "
"`inpaint_audio` + `inpaint_mask` (inpainting), not both.")
if init_audio is not None:
batch.extra["init_latent"] = self._encode_audio_reference(init_audio, device)
if inpaint_audio is not None and inpaint_mask is not None:
batch.extra["inpaint_reference_latent"] = self._encode_audio_reference(inpaint_audio, device)
batch.extra["inpaint_mask_latent"] = self._prepare_mask(inpaint_mask, latent_sample_size, device)
return batch
def _hop_length(self) -> int:
return int(self.vae.hop_length)
def _encode_audio_reference(self, audio, device: torch.device) -> torch.Tensor:
"""Pad/truncate to `sample_size` and encode via the VAE.
`audio` may be a tensor (`[samples]`, `[C, samples]`, or
`[B, C, samples]`) at the model's sample rate, or a path to any
audio-bearing file (`.wav` / `.mp3` / `.mp4` / `.m4a` / `.flac`,
...) that PyAV can decode — we resample on load so callers don't
have to.
"""
assert self.vae is not None, "VAE required for init_audio / inpaint_audio encoding"
# VAE may be CPU-parked under `vae_cpu_offload=True`.
self.vae = self.vae.to(device)
if isinstance(audio, str | os.PathLike):
audio = _decode_audio_file(audio, target_sr=self.sample_rate)
audio = audio.to(device=device, dtype=torch.float32)
if audio.dim() == 1:
audio = audio.unsqueeze(0).unsqueeze(0)
elif audio.dim() == 2:
audio = audio.unsqueeze(0)
# Match expected channel count (mono → repeat to stereo).
if audio.shape[1] == 1 and self.audio_channels == 2:
audio = audio.repeat(1, 2, 1)
elif audio.shape[1] == 2 and self.audio_channels == 1:
audio = audio.mean(dim=1, keepdim=True)
# Pad/truncate to model sample_size.
cur_len = audio.shape[-1]
if cur_len < self.sample_size:
audio = F.pad(audio, (0, self.sample_size - cur_len))
elif cur_len > self.sample_size:
audio = audio[..., :self.sample_size]
# Stochastic sample (the next random draw after the latent
# `randn` above), so encode-noise stays on the seeded sequence.
return self.vae.encode(audio.to(next(self.vae.parameters()).dtype)).sample()
def _prepare_mask(self, mask, latent_len: int, device: torch.device) -> torch.Tensor:
"""Pad/truncate a binary mask to `sample_size`, then
nearest-resample to `[1, 1, latent_len]`. Convention: 1 = keep
the reference, 0 = regenerate.
`mask` may be a `[samples]` tensor at the model sample rate or a
`(keep_seconds, total_seconds)` tuple — the tuple form builds
"keep first K seconds, regenerate the rest" automatically.
"""
if isinstance(mask, tuple) and len(mask) == 2:
keep_s, total_s = (float(x) for x in mask)
keep_n = int(keep_s * self.sample_rate)
total_n = int(total_s * self.sample_rate)
mask = torch.zeros(total_n, dtype=torch.float32)
mask[:keep_n] = 1.0
m = mask.to(device=device, dtype=torch.float32)
if m.dim() == 1:
m = m.unsqueeze(0)
cur_len = m.shape[-1]
if cur_len < self.sample_size:
m = F.pad(m, (0, self.sample_size - cur_len))
elif cur_len > self.sample_size:
m = m[..., :self.sample_size]
return F.interpolate(m.unsqueeze(1), size=latent_len, mode="nearest")
def _decode_audio_file(path, target_sr: int) -> torch.Tensor:
"""Decode any audio-bearing file (wav, mp3, mp4, m4a, flac, ...) via
PyAV and resample to `target_sr`. Returns `[channels, samples]`
float32 in roughly [-1, 1].
PyAV is already a FastVideo dep (used for muxing in
`VideoGenerator._mux_audio`). `torchaudio.load` on container
formats (mp4 / m4a) routes through `torchcodec`, which pulls in a
full CUDA NVRTC stack we don't otherwise need.
"""
import av
import numpy as np
container = av.open(str(path))
audio_stream = next(s for s in container.streams if s.type == "audio")
resampler = av.AudioResampler(format="fltp", layout="stereo", rate=target_sr)
chunks: list = []
for frame in container.decode(audio_stream):
for resampled in resampler.resample(frame):
chunks.append(resampled.to_ndarray())
for resampled in resampler.resample(None):
chunks.append(resampled.to_ndarray())
container.close()
if not chunks:
raise RuntimeError(f"No audio frames decoded from {path}")
waveform = np.concatenate(chunks, axis=-1)
if waveform.ndim == 1:
waveform = waveform[None, :]
return torch.from_numpy(waveform).float()
@@ -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
@@ -192,21 +190,6 @@ class ForwardBatch:
ltx2_stg_blocks_video: list[int] = field(default_factory=list)
ltx2_stg_blocks_audio: list[int] = field(default_factory=list)
# Stable Audio (T2A): clip start/end in seconds. Parallels the
# `SamplingParam` fields of the same name; the
# `StableAudioConditioningStage` / `DecodingStage` read them.
audio_start_in_s: float | None = None
audio_end_in_s: float | None = None
# Stable Audio A2A variation + inpainting payloads (parallel to
# `SamplingParam`). `Any` because we accept torch tensors or numpy
# arrays the user supplies; the latent-prep stage normalises shapes.
init_audio: Any = None
init_audio_strength: float | None = None
init_noise_level: float | None = None
inpaint_audio: Any = None
inpaint_mask: Any = None
n_tokens: int | None = None
# Other parameters that may be needed by specific schedulers
@@ -223,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)
@@ -1,193 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Preprocess Cosmos 2.5 overfit data into parquet format.
Encodes videos with the Cosmos (Wan-style) VAE and captions with the
Reason1 (Qwen2.5-VL) text encoder into the t2v parquet schema.
Usage:
CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
"""
import json
import os
import cv2
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.utils import maybe_download_model
# --- Config ---
NUM_FRAMES = 93 # 4*23+1 → 24 latent frames
MAX_HEIGHT = 480
MAX_WIDTH = 832
TRAIN_FPS = 16.0
DATA_DIR = "data/cosmos_overfit"
OUTPUT_DIR = "data/cosmos25_overfit_preprocessed"
MODEL_REPO = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
# The VAE is architecturally identical to Cosmos Predict2;
# use the Predict2 model for VAE since its weights are in
# standard diffusers format.
VAE_REPO = "nvidia/Cosmos-Predict2-2B-Video2World"
def load_video(path: str, num_frames: int) -> torch.Tensor:
"""Load video as [1, C, T, H, W] in [-1, 1]."""
cap = cv2.VideoCapture(path)
frames: list[np.ndarray] = []
while len(frames) < num_frames:
ret, frame = cap.read()
if not ret:
break
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
cap.release()
if len(frames) < num_frames:
while len(frames) < num_frames:
frames.append(frames[-1])
frames = frames[:num_frames]
video = np.stack(frames, axis=0)
video = torch.from_numpy(video).float()
video = video / 127.5 - 1.0 # [0,255] -> [-1,1]
video = video.permute(3, 0, 1, 2).unsqueeze(0) # [1,C,T,H,W]
return video
def main() -> None:
device = torch.device("cuda:0")
model_path = maybe_download_model(MODEL_REPO)
os.makedirs(OUTPUT_DIR, exist_ok=True)
# Load captions
with open(os.path.join(DATA_DIR, "videos2caption.json")) as f:
caption_data = json.load(f)
# --- Load VAE (Wan-style, same arch for Cosmos 2 and 2.5) ---
print("Loading Cosmos VAE (AutoencoderKLWan)...")
vae_path = maybe_download_model(VAE_REPO)
from diffusers import AutoencoderKLWan
vae = AutoencoderKLWan.from_pretrained(
vae_path,
subfolder="vae",
torch_dtype=torch.float16,
).to(device).eval()
print(f"VAE loaded "
f"({sum(p.numel() for p in vae.parameters())/1e6:.0f}M)")
# --- Load Reason1 (Qwen2.5-VL) text encoder ---
print("Loading Reason1 text encoder...")
from fastvideo.configs.pipelines.cosmos2_5 import (
Cosmos25Config, )
from fastvideo.models.encoders.reason1 import (
Reason1TextEncoder, )
pipeline_cfg = Cosmos25Config()
text_enc_cfg = pipeline_cfg.text_encoder_configs[0]
text_enc_path = os.path.join(model_path, "text_encoder")
# Instantiate Reason1TextEncoder with config and checkpoint
text_encoder = Reason1TextEncoder(
text_enc_cfg,
checkpoint_path=text_enc_path,
)
# Load weights from safetensors into the meta-device model.
# Materialize empty tensors in bf16 on the target device,
# then overwrite with checkpoint weights.
text_encoder = text_encoder.to_empty(device=device)
text_encoder = text_encoder.to(torch.bfloat16)
import glob
from safetensors.torch import load_file
sd: dict[str, torch.Tensor] = {}
for sf in sorted(glob.glob(os.path.join(text_enc_path, "*.safetensors"))):
sd.update(load_file(sf, device=str(device)))
sd = {k: v.to(torch.bfloat16) for k, v in sd.items()}
text_encoder.load_state_dict(sd, strict=False, assign=True)
del sd
torch.cuda.empty_cache()
text_encoder = text_encoder.eval()
print("Reason1 text encoder loaded")
# --- Process each video ---
records = []
for idx, item in enumerate(caption_data):
video_name = item["path"]
record_id = f"{idx:04d}_{video_name}"
caption = item["cap"][0]
video_path = os.path.join(DATA_DIR, "videos", video_name)
print(f"\nProcessing: {video_name}")
print(f" Caption: {caption[:80]}...")
# Encode video
video = load_video(video_path, NUM_FRAMES).to(device=device, dtype=torch.float16)
print(f" Video shape: {video.shape}")
with torch.no_grad():
latent_dist = vae.encode(video).latent_dist
latent = latent_dist.mean.squeeze(0).float().cpu()
print(f" Latent shape: {latent.shape}")
# Encode text with Reason1 (Qwen2.5-VL)
with torch.no_grad():
text_embedding = text_encoder.compute_text_embeddings(
[caption],
device=device,
)
text_embedding = text_embedding.squeeze(0).float().cpu()
print(f" Text embedding shape: {text_embedding.shape}")
record = {
"id": record_id,
"vae_latent_bytes": latent.numpy().tobytes(),
"vae_latent_shape": list(latent.shape),
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
"text_embedding_bytes": (text_embedding.numpy().tobytes()),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
"file_name": video_name,
"caption": caption,
"media_type": "video",
"width": MAX_WIDTH,
"height": MAX_HEIGHT,
"num_frames": NUM_FRAMES,
"duration_sec": NUM_FRAMES / TRAIN_FPS,
"fps": TRAIN_FPS,
}
records.append(record)
# Clean up
del text_encoder, vae
torch.cuda.empty_cache()
# Write parquet
table = pa.table(
{k: [r[k] for r in records]
for k in records[0]},
schema=pyarrow_schema_t2v,
)
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
pq.write_table(table, output_path)
print(f"\nWrote {len(records)} records to {output_path}")
# Write T2W validation prompts (no image_path for T2W)
val_prompts = {
"data": [{
"caption": item["cap"][0],
} for item in caption_data],
}
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
with open(val_path, "w") as f:
json.dump(val_prompts, f, indent=2)
print(f"Wrote validation prompts to {val_path}")
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
if __name__ == "__main__":
main()
@@ -1,196 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Preprocess Cosmos-Predict2 overfit data into parquet format.
Encodes videos with the Cosmos (Wan-style) VAE and captions with the
single T5 Large text encoder into the t2v parquet schema expected by
the training framework.
Usage:
CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos_overfit.py
"""
import json
import os
import cv2
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from fastvideo.configs.models.encoders import T5LargeConfig
from fastvideo.configs.models.encoders.base import BaseEncoderOutput
from fastvideo.configs.pipelines.cosmos import t5_large_postprocess_text
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.utils import maybe_download_model
# --- Config ---
NUM_FRAMES = 93 # 4*23+1 for temporal compression ratio 4 -> 24 latent frames
MAX_HEIGHT = 480
MAX_WIDTH = 832
TRAIN_FPS = 16.0
DATA_DIR = "data/cosmos_overfit"
OUTPUT_DIR = "data/cosmos_overfit_preprocessed"
MODEL_REPO = "nvidia/Cosmos-Predict2-2B-Video2World"
def load_video(path: str, num_frames: int) -> torch.Tensor:
"""Load video as [1, C, T, H, W] in [-1, 1]."""
cap = cv2.VideoCapture(path)
frames: list[np.ndarray] = []
while len(frames) < num_frames:
ret, frame = cap.read()
if not ret:
break
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
cap.release()
if len(frames) < num_frames:
# Repeat last frame to fill
while len(frames) < num_frames:
frames.append(frames[-1])
frames = frames[:num_frames]
video = np.stack(frames, axis=0)
video = torch.from_numpy(video).float()
video = video / 127.5 - 1.0 # [0,255] -> [-1,1]
video = video.permute(3, 0, 1, 2).unsqueeze(0) # [1,C,T,H,W]
return video
def main() -> None:
device = torch.device("cuda:0")
model_path = maybe_download_model(MODEL_REPO)
os.makedirs(OUTPUT_DIR, exist_ok=True)
# Load captions
with open(os.path.join(DATA_DIR, "videos2caption.json")) as f:
caption_data = json.load(f)
# --- Load VAE ---
# Cosmos-Predict2-2B ships a Wan-style VAE; vae/config.json declares
# `_class_name: AutoencoderKLWan`, so diffusers can load it directly.
print("Loading Cosmos VAE (AutoencoderKLWan)...")
from diffusers import AutoencoderKLWan
vae = AutoencoderKLWan.from_pretrained(
model_path,
subfolder="vae",
torch_dtype=torch.float16,
).to(device).eval()
print(f"VAE loaded ({sum(p.numel() for p in vae.parameters())/1e6:.0f}M)")
# --- Load T5 Large text encoder ---
print("Loading T5 Large text encoder...")
from transformers import AutoTokenizer, T5EncoderModel
t5_cfg = T5LargeConfig()
tok_kwargs = dict(t5_cfg.tokenizer_kwargs)
tokenizer = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer"))
text_encoder = T5EncoderModel.from_pretrained(
os.path.join(model_path, "text_encoder"),
torch_dtype=torch.bfloat16,
).to(device).eval()
# --- Process each video ---
records = []
for idx, item in enumerate(caption_data):
video_name = item["path"]
record_id = f"{idx:04d}_{video_name}"
caption = item["cap"][0]
video_path = os.path.join(DATA_DIR, "videos", video_name)
print(f"\nProcessing: {video_name}")
print(f" Caption: {caption[:80]}...")
# Encode video
video = load_video(video_path, NUM_FRAMES).to(device=device, dtype=torch.float16)
print(f" Video shape: {video.shape}")
with torch.no_grad():
latent_dist = vae.encode(video).latent_dist
# Cast to fp32 — dataloader hardcodes np.float32
latent = latent_dist.mean.squeeze(0).float().cpu()
print(f" Latent shape: {latent.shape}")
# Encode text with T5
with torch.no_grad():
inputs = tokenizer(caption, **tok_kwargs).to(device)
outputs = text_encoder(**inputs)
enc_out = BaseEncoderOutput(
last_hidden_state=outputs.last_hidden_state,
attention_mask=inputs["attention_mask"],
)
# [1, max_len, 1024], zeros beyond real length
t5_embed = t5_large_postprocess_text(enc_out).squeeze(0)
# Trim to real sequence length so dataloader's pad() builds
# the correct attention mask.
real_len = int(inputs["attention_mask"].sum().item())
text_embedding = t5_embed[:real_len].float().cpu() # [seq, 1024]
print(f" Text embedding shape: {text_embedding.shape}")
record = {
"id": record_id,
"vae_latent_bytes": latent.numpy().tobytes(),
"vae_latent_shape": list(latent.shape),
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
"text_embedding_bytes": text_embedding.numpy().tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
"file_name": video_name,
"caption": caption,
"media_type": "video",
"width": MAX_WIDTH,
"height": MAX_HEIGHT,
"num_frames": NUM_FRAMES,
"duration_sec": NUM_FRAMES / TRAIN_FPS,
"fps": TRAIN_FPS,
}
records.append(record)
# Clean up encoders
del text_encoder, tokenizer, vae
# Write parquet
table = pa.table(
{k: [r[k] for r in records]
for k in records[0]},
schema=pyarrow_schema_t2v,
)
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
pq.write_table(table, output_path)
print(f"\nWrote {len(records)} records to {output_path}")
# Extract first frame from first video as V2W conditioning image
import cv2
first_video = os.path.join(DATA_DIR, "videos", caption_data[0]["path"])
cap = cv2.VideoCapture(first_video)
ret, frame = cap.read()
cap.release()
cond_frame_path = os.path.join(OUTPUT_DIR, "cond_frame.png")
if ret:
cv2.imwrite(cond_frame_path, frame)
print(f"Saved conditioning frame to {cond_frame_path}")
# Write validation prompts for callback
# Wrap in "data" key — ValidationDataset expects field="data"
# Use "caption" field — ValidationDataset aliases it to "prompt"
# Include image_path for V2W conditioning during validation
val_prompts = {
"data": [{
"caption": item["cap"][0],
"image_path": "cond_frame.png",
} for item in caption_data]
}
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
with open(val_path, "w") as f:
json.dump(val_prompts, f, indent=2)
print(f"Wrote validation prompts to {val_path}")
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
if __name__ == "__main__":
main()
+166 -189
View File
@@ -518,42 +518,13 @@ class DenoisingStage(PipelineStage):
class CosmosDenoisingStage(DenoisingStage):
"""Denoising stage for Cosmos models.
Uses FlowMatchEulerDiscreteScheduler with manual EDM
preconditioning (c_in, c_skip, c_out) to match the
pretrained Cosmos model's training convention.
"""
Denoising stage for Cosmos models using FlowMatchEulerDiscreteScheduler.
"""
def __init__(self, transformer, scheduler, pipeline=None) -> None:
super().__init__(transformer, scheduler, pipeline)
def _run_transformer(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
condition_mask: torch.Tensor,
padding_mask: torch.Tensor,
target_dtype: torch.dtype,
step_index: int,
batch: ForwardBatch,
) -> torch.Tensor:
with set_forward_context(
current_timestep=step_index,
attn_metadata=None,
forward_batch=batch,
):
return self.transformer(
hidden_states=hidden_states.to(target_dtype),
timestep=timestep.to(target_dtype),
encoder_hidden_states=encoder_hidden_states.to(target_dtype),
fps=24,
condition_mask=condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
def forward(
self,
batch: ForwardBatch,
@@ -562,188 +533,199 @@ class CosmosDenoisingStage(DenoisingStage):
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"],
fastvideo_args,
)
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
if hasattr(self.transformer, "module"):
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
if hasattr(self.transformer, 'module'):
transformer_dtype = next(self.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32 and not fastvideo_args.disable_autocast)
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
latents = batch.latents
num_inference_steps = batch.num_inference_steps
guidance_scale = batch.guidance_scale
do_cfg = (batch.do_classifier_free_guidance and batch.negative_prompt_embeds is not None)
sigma_data = float(getattr(self.scheduler.config, "sigma_data", 1.0))
sigma_max = 80.0
sigma_min = 0.002
sigma_data = 1.0
final_sigmas_type = "sigma_min"
self.scheduler.set_timesteps(
num_inference_steps,
device=latents.device,
)
if self.scheduler is not None:
self.scheduler.register_to_config(
sigma_max=sigma_max,
sigma_min=sigma_min,
sigma_data=sigma_data,
final_sigmas_type=final_sigmas_type,
)
self.scheduler.set_timesteps(num_inference_steps, device=latents.device)
timesteps = self.scheduler.timesteps
# Clamp terminal sigma to sigma_min (avoid zero).
if (hasattr(self.scheduler.config, "final_sigmas_type")
if (hasattr(self.scheduler.config, 'final_sigmas_type')
and self.scheduler.config.final_sigmas_type == "sigma_min" and len(self.scheduler.sigmas) > 1):
self.scheduler.sigmas[-1] = self.scheduler.sigmas[-2]
conditioning_latents = getattr(
batch,
"conditioning_latents",
None,
)
cond_indicator = getattr(batch, "cond_indicator", None)
uncond_indicator = getattr(
batch,
"uncond_indicator",
None,
)
conditioning_latents = getattr(batch, 'conditioning_latents', None)
unconditioning_latents = conditioning_latents
augment_sigma = torch.tensor(
[0.001],
device=latents.device,
dtype=torch.float32,
)
padding_mask = torch.zeros(
1,
1,
batch.height,
batch.width,
device=latents.device,
dtype=target_dtype,
)
condition_mask = (batch.cond_mask.to(target_dtype)
if hasattr(batch, "cond_mask") and batch.cond_mask is not None else None)
uncond_condition_mask = (batch.uncond_mask.to(target_dtype)
if hasattr(batch, "uncond_mask") and batch.uncond_mask is not None else condition_mask)
if condition_mask is None:
b, c, tf, h, w = latents.shape
condition_mask = torch.zeros(
b,
1,
tf,
h,
w,
device=latents.device,
dtype=target_dtype,
)
uncond_condition_mask = condition_mask
with self.progress_bar(total=num_inference_steps, ) as progress_bar:
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if hasattr(self, "interrupt") and self.interrupt:
if hasattr(self, 'interrupt') and self.interrupt:
continue
sigma = self.scheduler.sigmas[i]
is_aug_greater = bool(augment_sigma >= sigma)
current_sigma = self.scheduler.sigmas[i]
current_t = current_sigma / (current_sigma + 1)
c_in = 1 - current_t
c_skip = 1 - current_t
c_out = -current_t
# EDM preconditioning coefficients.
c_in = 1.0 / (sigma**2 + sigma_data**2)**0.5
c_in_aug = 1.0 / (augment_sigma**2 + sigma_data**2)**0.5
c_skip = sigma_data**2 / (sigma**2 + sigma_data**2)
c_out = (sigma * sigma_data / (sigma**2 + sigma_data**2)**0.5)
timestep = current_t.view(1, 1, 1, 1, 1).expand(latents.size(0), -1, latents.size(2), -1,
-1) # [B, 1, T, 1, 1]
# The model expects timestep = sigma * 1000
# (FlowMatchEulerDiscreteScheduler convention).
timestep_expanded = t.expand(latents.shape[0], ).to(target_dtype)
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
with torch.autocast(
device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled,
):
# --- Conditioning frame injection ---
cur_ci = (cond_indicator * 0 if cond_indicator is not None and is_aug_greater else cond_indicator)
cond_latent = latents * c_in
cond_latent = latents.clone()
if (cur_ci is not None and conditioning_latents is not None):
cn = torch.randn_like(
latents,
dtype=torch.float32,
)
cf = (conditioning_latents + cn * augment_sigma[:, None, None, None, None])
cf = cf * c_in_aug / c_in
cond_latent = (cur_ci * cf + (1 - cur_ci) * cond_latent)
# Manual EDM input scaling.
model_input = cond_latent * c_in
noise_pred_cond = self._run_transformer(
model_input,
timestep_expanded,
batch.prompt_embeds[0],
condition_mask,
padding_mask,
target_dtype,
i,
batch,
)
# EDM output → x0 prediction.
cond_x0 = (c_skip * latents + c_out * noise_pred_cond.float())
if (cur_ci is not None and conditioning_latents is not None):
cond_x0 = (cur_ci * conditioning_latents + (1 - cur_ci) * cond_x0)
# --- CFG: unconditional pass ---
if do_cfg:
cur_ui = (uncond_indicator *
0 if uncond_indicator is not None and is_aug_greater else uncond_indicator)
uncond_latent = latents.clone()
if (cur_ui is not None and conditioning_latents is not None):
un = torch.randn_like(
latents,
dtype=torch.float32,
)
uf = (conditioning_latents + un * augment_sigma[:, None, None, None, None])
uf = uf * c_in_aug / c_in
uncond_latent = (cur_ui * uf + (1 - cur_ui) * uncond_latent)
uncond_input = uncond_latent * c_in
noise_pred_uncond = (self._run_transformer(
uncond_input,
timestep_expanded,
batch.negative_prompt_embeds[0],
uncond_condition_mask,
padding_mask,
target_dtype,
i,
if hasattr(
batch,
))
uncond_x0 = (c_skip * latents + c_out * noise_pred_uncond.float())
if (cur_ui is not None and conditioning_latents is not None):
uncond_x0 = (cur_ui * conditioning_latents + (1 - cur_ui) * uncond_x0)
final_x0 = (cond_x0 + guidance_scale * (cond_x0 - uncond_x0))
'cond_indicator') and batch.cond_indicator is not None and conditioning_latents is not None:
cond_latent = batch.cond_indicator * conditioning_latents + (1 -
batch.cond_indicator) * cond_latent
else:
final_x0 = cond_x0
logger.warning(
"Step %s: Missing conditioning data - cond_indicator: %s, conditioning_latents: %s", i,
hasattr(batch, 'cond_indicator'), conditioning_latents is not None)
# Convert x0 to velocity for
# FlowMatchEulerDiscreteScheduler.
velocity = (latents - final_x0) / sigma.clamp(min=1e-6)
cond_latent = cond_latent.to(target_dtype)
latents = self.scheduler.step(
velocity,
t,
latents,
return_dict=False,
)[0]
cond_timestep = timestep
if hasattr(batch, 'cond_indicator') and batch.cond_indicator is not None:
sigma_conditioning = 0.0001
t_conditioning = sigma_conditioning / (sigma_conditioning + 1)
cond_timestep = batch.cond_indicator * t_conditioning + (1 - batch.cond_indicator) * timestep
cond_timestep = cond_timestep.to(target_dtype)
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
# Use conditioning masks from CosmosLatentPreparationStage
condition_mask = batch.cond_mask.to(target_dtype) if hasattr(batch, 'cond_mask') else None
padding_mask = torch.zeros(1,
1,
batch.height,
batch.width,
device=cond_latent.device,
dtype=target_dtype)
# Fallback if masks not available
if condition_mask is None:
batch_size, num_channels, num_frames, height, width = cond_latent.shape
condition_mask = torch.zeros(batch_size,
1,
num_frames,
height,
width,
device=cond_latent.device,
dtype=target_dtype)
noise_pred = self.transformer(
hidden_states=cond_latent,
timestep=cond_timestep.to(target_dtype),
encoder_hidden_states=batch.prompt_embeds[0].to(target_dtype),
fps=24, # TODO: get fps from batch or config
condition_mask=condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
cond_pred = (c_skip * latents + c_out * noise_pred.float()).to(target_dtype)
if hasattr(
batch,
'cond_indicator') and batch.cond_indicator is not None and conditioning_latents is not None:
cond_pred = batch.cond_indicator * conditioning_latents + (1 - batch.cond_indicator) * cond_pred
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds is not None:
uncond_latent = latents * c_in
if hasattr(batch, 'uncond_indicator'
) and batch.uncond_indicator is not None and unconditioning_latents is not None:
uncond_latent = batch.uncond_indicator * unconditioning_latents + (
1 - batch.uncond_indicator) * uncond_latent
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
uncond_condition_mask = batch.uncond_mask.to(target_dtype) if hasattr(
batch, 'uncond_mask') and batch.uncond_mask is not None else condition_mask
uncond_timestep = timestep
if hasattr(batch, 'uncond_indicator') and batch.uncond_indicator is not None:
sigma_conditioning = 0.0001
t_conditioning = sigma_conditioning / (sigma_conditioning + 1)
uncond_timestep = batch.uncond_indicator * t_conditioning + (
1 - batch.uncond_indicator) * timestep
uncond_timestep = uncond_timestep.to(target_dtype)
noise_pred_uncond = self.transformer(
hidden_states=uncond_latent.to(target_dtype),
timestep=uncond_timestep.to(target_dtype),
encoder_hidden_states=batch.negative_prompt_embeds[0].to(target_dtype),
fps=24, # TODO: get fps from batch or config
condition_mask=uncond_condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
uncond_pred = (c_skip * latents + c_out * noise_pred_uncond.float()).to(target_dtype)
if hasattr(batch, 'uncond_indicator'
) and batch.uncond_indicator is not None and unconditioning_latents is not None:
uncond_pred = batch.uncond_indicator * unconditioning_latents + (
1 - batch.uncond_indicator) * uncond_pred
guidance_diff = cond_pred - uncond_pred
final_pred = cond_pred + guidance_scale * guidance_diff
else:
final_pred = cond_pred
# Convert to noise for scheduler step
if current_sigma > 1e-8:
noise_for_scheduler = (latents - final_pred) / current_sigma
else:
logger.warning("Step %s: current_sigma too small (%s), using final_pred directly", i, current_sigma)
noise_for_scheduler = final_pred
if torch.isnan(noise_for_scheduler).sum() > 0:
logger.error("Step %s: NaN detected in noise_for_scheduler, sum: %s", i,
noise_for_scheduler.float().sum().item())
logger.error("Step %s: latents sum: %s, final_pred sum: %s, current_sigma: %s", i,
latents.float().sum().item(),
final_pred.float().sum().item(), current_sigma)
latents = self.scheduler.step(noise_for_scheduler, t, latents, **extra_step_kwargs,
return_dict=False)[0]
progress_bar.update()
batch.latents = latents
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
@@ -789,21 +771,16 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
},
)
# Detect the actual weight dtype. FSDP-wrapped models may
# report fp32 via next(parameters()) even when the physical
# weights are bf16. Walk through parameters to find one
# that is NOT fp32 (the real checkpoint dtype).
target_dtype = torch.bfloat16 # safe default for Cosmos 2.5
for p in self.transformer.parameters():
if p.dtype != torch.float32:
target_dtype = p.dtype
break
if hasattr(self.transformer, 'module'):
transformer_dtype = next(self.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
latents = batch.latents
if latents is None:
raise ValueError("latents must be provided for "
"Cosmos25DenoisingStage")
raise ValueError("latents must be provided for Cosmos25DenoisingStage")
guidance_scale = batch.guidance_scale
if batch.timesteps is None:
@@ -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:
-45
View File
@@ -48,7 +48,6 @@ from fastvideo.configs.pipelines.wan import (
WanT2V720PConfig,
)
from fastvideo.configs.pipelines.sd35 import SD35Config
from fastvideo.configs.pipelines.stable_audio import (StableAudioOpenSmallConfig, StableAudioT2AConfig)
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import WorkloadType
@@ -243,47 +242,6 @@ def _register_configs() -> None:
default_preset="ltx2_distilled",
)
# Stable Audio Open (text-to-audio). Both variants must be loaded
# from the FastVideo-curated converted Diffusers-format repos —
# the upstream `stabilityai/stable-audio-open-{1.0,small}` repos
# ship `model.safetensors` as a single monolithic checkpoint with
# no per-component subfolders our standard loader can consume. See
# `scripts/checkpoint_conversion/stable_audio_to_diffusers.py`.
# NOTE: WorkloadType has no T2A variant yet (REVIEW item 28); using
# T2V as the placeholder until the enum is extended.
register_configs(
sampling_param_cls=None,
pipeline_config_cls=StableAudioT2AConfig,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"FastVideo/stable-audio-open-1.0-Diffusers",
],
# Substring match against HF cache snapshot paths (the lookup
# runs on the resolved local directory, which uses `--` between
# org and repo: `models--FastVideo--stable-audio-open-1.0-Diffusers`).
model_detectors=[
lambda path: "stable-audio-open-1" in path.lower(),
],
model_family="stable_audio",
default_preset="stable_audio_open_1_0_base",
)
# Small variant uses its own `pipeline_config_cls` so it picks up
# the smaller (524288-sample) training window in `sample_size` /
# `max_audio_duration_s`.
register_configs(
sampling_param_cls=None,
pipeline_config_cls=StableAudioOpenSmallConfig,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"FastVideo/stable-audio-open-small-Diffusers",
],
model_detectors=[
lambda path: "stable-audio-open-small" in path.lower(),
],
model_family="stable_audio",
default_preset="stable_audio_open_small",
)
# Hunyuan 1.5 (specific)
register_configs(
sampling_param_cls=None,
@@ -828,8 +786,6 @@ def _register_presets() -> None:
ALL_PRESETS as MATRIXGAME_PRESETS, )
from fastvideo.pipelines.basic.sd35.presets import (
ALL_PRESETS as SD35_PRESETS, )
from fastvideo.pipelines.basic.stable_audio.presets import (
ALL_PRESETS as STABLE_AUDIO_PRESETS, )
from fastvideo.pipelines.basic.turbodiffusion.presets import (
ALL_PRESETS as TURBODIFFUSION_PRESETS, )
from fastvideo.pipelines.basic.wan.presets import (
@@ -847,7 +803,6 @@ def _register_presets() -> None:
LTX2_PRESETS,
MATRIXGAME_PRESETS,
SD35_PRESETS,
STABLE_AUDIO_PRESETS,
TURBODIFFUSION_PRESETS,
WAN_PRESETS,
)
+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())
+6 -26
View File
@@ -24,10 +24,6 @@ image = (modal.Image.from_registry(
os.environ.get("BUILDKITE_COMMIT", ""),
"BUILDKITE_PULL_REQUEST":
os.environ.get("BUILDKITE_PULL_REQUEST", ""),
"BUILDKITE_BRANCH":
os.environ.get("BUILDKITE_BRANCH", ""),
"TEST_SCOPE":
os.environ.get("TEST_SCOPE", ""),
"IMAGE_VERSION":
os.environ.get("IMAGE_VERSION", ""),
}))
@@ -70,21 +66,13 @@ def run_test(pytest_command: str):
{pytest_command}
"""
# result = subprocess.run(["/bin/bash", "-c", command],
# stdout=sys.stdout,
# stderr=sys.stderr,
# check=False)
# sys.exit(result.returncode)
result = subprocess.run(["/bin/bash", "-c", command],
stdout=sys.stdout,
stderr=sys.stderr,
check=False)
if result.returncode != 0:
raise RuntimeError(f"Test command failed with exit code {result.returncode}")
# On success, just return — don't call sys.exit()
sys.exit(result.returncode)
@app.function(gpu="H100:1",
image=image,
@@ -218,7 +206,7 @@ def run_self_forcing_tests():
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test(
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
)
@@ -240,20 +228,12 @@ def run_lora_extraction_tests():
timeout=1800,
secrets=[
modal.Secret.from_dict(
{"HF_API_KEY": os.environ.get("HF_API_KEY", ""),
"HF_REPO_ID": "FastVideo/performance-tracking"})
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
],
volumes={
"/root/data": model_vol,
})
volumes={"/root/data": model_vol})
def run_performance_tests():
run_test(
"export HF_HOME='/root/data/.cache' && "
"export PERFORMANCE_TRACKING_ROOT='/tmp/perf-tracking' && "
"hf auth login --token $HF_API_KEY && "
"pytest ./fastvideo/tests/performance -vs && "
"python ./fastvideo/tests/performance/compare_baseline.py && "
"python ./fastvideo/tests/performance/dashboard.py"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/performance -vs"
)
@@ -1,320 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Track performance results and compare against historical baseline.
This script:
1) reads current benchmark results from fastvideo/tests/performance/results,
2) writes normalized tracking records to the Modal volume path,
3) compares each current record against the mean of up to 5 prior records,
4) exits non-zero if any metric regresses by more than 15%.
"""
import glob
import json
import os
import re
import statistics
import sys
from huggingface_hub import HfApi, snapshot_download
from datetime import datetime, timezone
from typing import Any
from hf_store import sync_from_hf, upload_record, load_records_for_model, sanitize, safe_float
# Use the env var passed by Modal, fallback to a default if needed
HF_REPO_ID = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
HF_TOKEN = os.environ.get("HF_API_KEY")
RESULTS_DIR = os.path.join(
os.path.dirname(os.path.abspath(__file__)),
"results",
)
TRACKING_ROOT = os.environ.get(
"PERFORMANCE_TRACKING_ROOT",
"/tmp/perf-tracking",
)
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
def _should_persist_tracking() -> bool:
# test_scope = os.environ.get("TEST_SCOPE", "")
# branch = os.environ.get("BUILDKITE_BRANCH", "")
# return test_scope == "full" and branch == "main"
return True # only for testing purpose.
def _sanitize(value: str) -> str:
return re.sub(r"[^A-Za-z0-9._-]", "_", value)
def _safe_float(value: Any) -> float | None:
if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None
def _load_current_results() -> list[dict[str, Any]]:
pattern = os.path.join(RESULTS_DIR, "perf_*.json")
records: list[dict[str, Any]] = []
for path in sorted(glob.glob(pattern)):
with open(path, encoding="utf-8") as f:
records.append(json.load(f))
return records
def _normalize_record(result: dict[str, Any]) -> dict[str, Any]:
benchmark_id = result.get("benchmark_id", "unknown")
model_id = benchmark_id
timestamp = result.get("timestamp")
if not timestamp:
timestamp = datetime.now(timezone.utc).isoformat()
commit_sha = result.get("commit") or os.environ.get("BUILDKITE_COMMIT", "")
latency = _safe_float(result.get("avg_generation_time_s"))
throughput = _safe_float(result.get("throughput_fps"))
memory = _safe_float(result.get("max_peak_memory_mb"))
return {
"model_id": model_id,
"timestamp": timestamp,
"commit_sha": commit_sha,
"gpu_type": result.get("device", "unknown"),
"latency": latency,
"throughput": throughput,
"memory": memory,
"success": True,
}
def _write_tracking_record(record: dict[str, Any]) -> str:
model_dir = os.path.join(TRACKING_ROOT, _sanitize(record["model_id"]))
os.makedirs(model_dir, exist_ok=True)
timestamp = _sanitize(record["timestamp"])
commit = _sanitize(record["commit_sha"] or "unknown")
out_path = os.path.join(model_dir, f"{timestamp}_{commit}.json")
with open(out_path, "w", encoding="utf-8") as f:
json.dump(record, f, indent=2)
return out_path
def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
values = [
_safe_float(r.get(key))
for r in records
]
values = [v for v in values if v is not None]
if not values:
return None
return statistics.median(values)
def _check_regressions(
current: dict[str, Any],
baseline_records: list[dict[str, Any]],
max_regression: float,
) -> list[str]:
failures: list[str] = []
for metric in ("latency", "memory"):
baseline = _baseline_metric(baseline_records, metric)
curr = _safe_float(current.get(metric))
if baseline is None or curr is None or baseline <= 0:
continue
regression = (curr - baseline) / baseline
if regression > max_regression:
failures.append(
f"{current['model_id']} {metric} regressed by {regression * 100:.1f}% "
f"(current={curr:.3f}, baseline_median={baseline:.3f})"
)
baseline_tp = _baseline_metric(baseline_records, "throughput")
curr_tp = _safe_float(current.get("throughput"))
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
regression = (baseline_tp - curr_tp) / baseline_tp
if regression > max_regression:
failures.append(
f"{current['model_id']} throughput regressed by {regression * 100:.1f}% "
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})"
)
return failures
def _metric_delta_percent(
metric: str,
current: dict[str, Any],
baseline_records: list[dict[str, Any]],
) -> float | None:
curr = _safe_float(current.get(metric))
baseline = _baseline_metric(baseline_records, metric)
if curr is None or baseline is None or baseline <= 0:
return None
if metric in ("latency", "memory"):
return (curr - baseline) / baseline * 100.0
if metric == "throughput":
return (baseline - curr) / baseline * 100.0
return None
def _compact_value(value: float | None, precision: int = 3) -> str:
if value is None:
return "n/a"
return f"{value:.{precision}f}"
def _build_summary_row(
record: dict[str, Any],
baseline_records: list[dict[str, Any]],
has_failed: bool
) -> dict[str, Any]:
"""Formats a single benchmark result into a row for the Markdown summary table."""
latency_base = _safe_float(_baseline_metric(baseline_records, "latency"))
throughput_base = _safe_float(_baseline_metric(baseline_records, "throughput"))
memory_base = _safe_float(_baseline_metric(baseline_records, "memory"))
# Calculate percentages for the 'Worst Regression' column
latency_reg = _metric_delta_percent("latency", record, baseline_records)
throughput_reg = _metric_delta_percent("throughput", record, baseline_records)
memory_reg = _metric_delta_percent("memory", record, baseline_records)
regressions = [v for v in (latency_reg, throughput_reg, memory_reg) if v is not None]
worst_regression_pct = max(regressions) if regressions else None
return {
"model_id": record["model_id"],
"gpu_type": record["gpu_type"],
"baseline_n": len(baseline_records),
"latency_curr": _safe_float(record.get("latency")),
"latency_base": latency_base,
"throughput_curr": _safe_float(record.get("throughput")),
"throughput_base": throughput_base,
"memory_curr": _safe_float(record.get("memory")),
"memory_base": memory_base,
"worst_regression_pct": worst_regression_pct,
"failed": has_failed,
}
def _build_markdown_summary(
summary_rows: list[dict[str, Any]],
max_regression: float,
) -> str:
lines = [
"## Performance Baseline Comparison",
"",
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
"",
"| Model | GPU | Baseline N | Latency (curr/base) | Throughput (curr/base) | Memory (curr/base) | Worst Regression | Status |",
"|---|---|---:|---|---|---|---:|---|",
]
for row in summary_rows:
latency = f"{_compact_value(row['latency_curr'])} / {_compact_value(row['latency_base'])}"
throughput = f"{_compact_value(row['throughput_curr'])} / {_compact_value(row['throughput_base'])}"
memory = f"{_compact_value(row['memory_curr'], 1)} / {_compact_value(row['memory_base'], 1)}"
worst_reg = "n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%"
status = "FAIL" if row["failed"] else "PASS"
lines.append(
f"| {row['model_id']} | {row['gpu_type']} | {row['baseline_n']} | "
f"{latency} | {throughput} | {memory} | {worst_reg} | {status} |"
)
return "\n".join(lines) + "\n"
def _emit_markdown_summary(markdown: str, commit_sha: str) -> None:
print("\n" + markdown)
# 1. Existing GitHub logic (safe to keep)
summary_path = os.environ.get("GITHUB_STEP_SUMMARY")
if summary_path:
with open(summary_path, "a", encoding="utf-8") as f:
f.write(markdown + "\n")
# 2. Write to Modal volume for Buildkite to pick up in post-run hook
try:
perf_reports_dir = "/root/data/perf_reports"
os.makedirs(perf_reports_dir, exist_ok=True)
short_sha = commit_sha[:7] if commit_sha else "unknown"
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
report_path = os.path.join(perf_reports_dir, f"perf_{short_sha}_{timestamp}.md")
with open(report_path, "w", encoding="utf-8") as f:
f.write(markdown + "\n")
print(f"Performance report written to {report_path}")
except Exception as e:
print(f"Failed to write performance report to Modal volume: {e}")
def main() -> int:
# Pull the current state of the world from HF
sync_from_hf(TRACKING_ROOT)
current_results = _load_current_results()
if not current_results:
print(f"No performance result files found in {RESULTS_DIR}")
return 0
all_failures = []
summary_rows = []
persist_tracking = _should_persist_tracking()
if persist_tracking:
print("Tracking persistence enabled: full-suite run on main branch")
else:
print("Tracking persistence disabled: only full-suite runs on main branch are persisted")
for raw in current_results:
record = _normalize_record(raw)
baseline_records = load_records_for_model(
TRACKING_ROOT, record["model_id"], record["gpu_type"],
last_n=5, successful_only=True
)
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
# Tag the current record based on the failure.
if not baseline_records:
# INITIALIZATION CASE: First run for this model/GPU
print(f"No baseline for {record['model_id']} on {record['gpu_type']}. Initializing...")
failures = []
record["success"] = True # The first run is always "successful"
else:
# COMPARISON CASE: Compare against the mean of the last 5 good runs
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
if failures:
record["success"] = False
all_failures.extend(failures)
else:
record["success"] = True
# 5. Persist to HF if we are on main
if persist_tracking:
# This writes the JSON with the "success" field to /tmp
current_path = _write_tracking_record(record)
# This pushes it to the FastVideo/performance-tracking repo
upload_record(current_path, record)
summary_row = _build_summary_row(record, baseline_records, bool(failures))
summary_rows.append(summary_row)
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
_emit_markdown_summary(markdown, commit_sha)
if all_failures:
print("Performance regression check failed:")
for item in all_failures:
print(f" - {item}")
return 1
print("Performance baseline comparison passed")
return 0
if __name__ == "__main__":
sys.exit(main())
-102
View File
@@ -1,102 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import os
import shutil
import subprocess
from datetime import datetime
import plotly.express as px
import pandas as pd
from hf_store import sync_from_hf, load_as_dataframe
# -----------------------------
# 1. Grouping
# -----------------------------
def group_data(df: pd.DataFrame):
# Group only by model+GPU so each group produces a time-series line.
# config_id (commit SHA) is carried as a column for hover/color use.
keys = ["model_id", "gpu_type"]
return df.groupby(keys, dropna=False)
# -----------------------------
# 2. Plot builder
# -----------------------------
def build_plots(df: pd.DataFrame) -> list:
figs = []
for (model_id, gpu_type), g in group_data(df):
g = g.sort_values("timestamp")
# One chart per metric so the y-axes aren't on wildly different scales
for metric in ("latency", "throughput", "memory"):
if g[metric].isna().all():
continue
fig = px.line(
g,
x="timestamp",
y=metric,
markers=True,
hover_data=["config_id", "commit_sha"],
title=f"{model_id} | {gpu_type} | {metric}",
labels={"timestamp": "Time", metric: metric},
)
figs.append(fig)
return figs
# -----------------------------
# 3. Render HTML dashboard
# -----------------------------
def render_html(figs: list, days: int) -> str:
html_parts = [
"<html>",
"<head><meta charset='utf-8'>",
"<style>body { font-family: sans-serif; margin: 2rem; }</style>",
"</head><body>",
f"<h2>Performance Dashboard (last {days} days)</h2>",
]
for fig in figs:
html_parts.append(fig.to_html(full_html=False, include_plotlyjs="cdn"))
html_parts.append("</body></html>")
return "\n".join(html_parts)
# -----------------------------
# 5. Main
# -----------------------------
def main() -> None:
days = int(os.environ.get("DASHBOARD_DAYS", "30"))
local_dir = sync_from_hf("/tmp/perf-tracking")
df = load_as_dataframe(local_dir, days=days)
if df.empty:
print("No data found")
return
# Sanity-check: log what we actually loaded
print(f"Loaded {len(df)} records across {df['model_id'].nunique()} model(s), "
f"{df['gpu_type'].nunique()} GPU type(s), "
f"date range: {df['timestamp'].min()} → {df['timestamp'].max()}")
figs = build_plots(df)
html = render_html(figs, days)
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
report_dir = "/root/data/perf_reports"
os.makedirs(report_dir, exist_ok=True)
filename = f"dashboard_{commit_sha}_{timestamp}.html"
output_file = os.path.join(report_dir, filename)
with open(output_file, "w", encoding="utf-8") as f:
f.write(html)
print(f"Dashboard generated: {output_file}")
if __name__ == "__main__":
main()
-254
View File
@@ -1,254 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared HuggingFace storage utilities for performance tracking.
Provides a single place for:
- Syncing the HF dataset repo to a local directory
- Loading raw JSON records (with optional recency filter)
- Loading records as a normalized pandas DataFrame
- Uploading individual result files back to HF
- Common helpers: sanitize, safe_float
"""
import glob
import json
import os
import re
from datetime import datetime, timedelta, timezone
from typing import Any
import pandas as pd
from huggingface_hub import HfApi, snapshot_download
# ---------------------------------------------------------------------------
# Configuration — read once at import time, shared across both consumers
# ---------------------------------------------------------------------------
HF_REPO_ID: str = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
HF_TOKEN: str | None = os.environ.get("HF_API_KEY")
# ---------------------------------------------------------------------------
# Low-level helpers
# ---------------------------------------------------------------------------
def sanitize(value: str) -> str:
"""Return a filesystem- and HF-path-safe version of *value*."""
return re.sub(r"[^A-Za-z0-9._-]", "_", value)
def safe_float(value: Any) -> float | None:
"""Coerce *value* to float, returning None on failure."""
if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None
# ---------------------------------------------------------------------------
# HF I/O
# ---------------------------------------------------------------------------
def sync_from_hf(local_dir: str) -> str:
"""Download the HF dataset repo snapshot to *local_dir*.
Returns *local_dir* so callers can chain: ``load_records(sync_from_hf(...))``.
On failure (empty repo, no credentials, network error) the function logs a
warning and returns *local_dir* unchanged so the caller can still work with
whatever is already on disk.
"""
if not HF_REPO_ID:
print("hf_store: HF_REPO_ID not set, skipping sync.")
return local_dir
print(f"hf_store: syncing from {HF_REPO_ID} → {local_dir}")
try:
snapshot_download(
repo_id=HF_REPO_ID,
repo_type="dataset",
local_dir=local_dir,
token=HF_TOKEN,
allow_patterns="*.json",
)
except Exception as exc:
print(f"hf_store: sync skipped — {exc}")
return local_dir
def upload_record(local_path: str, record: dict[str, Any]) -> None:
"""Upload *local_path* to the HF repo under ``<model_id>/<filename>``.
Silently skips if HF_TOKEN is absent so local/CI runs without credentials
don't crash.
"""
if not HF_TOKEN:
print("hf_store: HF_API_KEY not set, skipping upload.")
return
model_id = record.get("model_id", "unknown")
path_in_repo = f"{sanitize(model_id)}/{os.path.basename(local_path)}"
commit_sha = (record.get("commit_sha") or "unknown")[:7]
api = HfApi(token=HF_TOKEN)
try:
api.upload_file(
path_or_fileobj=local_path,
path_in_repo=path_in_repo,
repo_id=HF_REPO_ID,
repo_type="dataset",
commit_message=f"Perf: {model_id} at {commit_sha}",
)
print(f"hf_store: uploaded → {HF_REPO_ID}/{path_in_repo}")
except Exception as exc:
print(f"hf_store: upload failed — {exc}")
# ---------------------------------------------------------------------------
# Record loading
# ---------------------------------------------------------------------------
def load_records(
local_dir: str,
*,
days: int | None = None,
successful_only: bool = False,
) -> list[dict[str, Any]]:
"""Return raw JSON dicts from *local_dir*.
Args:
local_dir: Root directory previously populated by :func:`sync_from_hf`.
days: When set, discard records whose ``timestamp`` is older than this
many days. Records with a missing/unparseable timestamp are kept.
successful_only: When True, only records with ``success=True`` are
returned. Useful when building a regression baseline.
Returns:
List of raw dicts sorted by ``timestamp`` ascending (records that could
not be parsed are silently skipped).
"""
cutoff: datetime | None = None
if days is not None:
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
records: list[dict[str, Any]] = []
for path in sorted(glob.glob(os.path.join(local_dir, "**", "*.json"), recursive=True)):
try:
with open(path, encoding="utf-8") as fh:
data: dict[str, Any] = json.load(fh)
except (OSError, json.JSONDecodeError):
continue
if successful_only and not data.get("success", True):
continue
if cutoff is not None:
raw_ts = data.get("timestamp")
if raw_ts:
try:
ts = datetime.fromisoformat(raw_ts)
if ts.tzinfo is None:
ts = ts.replace(tzinfo=timezone.utc)
if ts < cutoff:
continue
except ValueError:
pass # keep records with unparseable timestamps
records.append(data)
return records
def load_records_for_model(
local_dir: str,
model_id: str,
gpu_type: str | None = None,
*,
last_n: int | None = None,
successful_only: bool = True,
) -> list[dict[str, Any]]:
"""Return records for a specific *model_id*, optionally filtered by GPU.
Args:
local_dir: Root directory previously populated by :func:`sync_from_hf`.
model_id: Matches the ``model_id`` field inside each JSON record.
gpu_type: When set, only records whose ``gpu_type`` matches are returned.
last_n: When set, return only the most recent *n* records (after all
other filters). Useful for sliding-window baseline calculations.
successful_only: Passed through to :func:`load_records`.
Returns:
List of matching dicts sorted by timestamp ascending.
"""
model_dir = os.path.join(local_dir, sanitize(model_id))
if not os.path.isdir(model_dir):
return []
records = load_records(model_dir, successful_only=successful_only)
if gpu_type is not None:
records = [r for r in records if r.get("gpu_type") == gpu_type]
if last_n is not None:
records = records[-last_n:]
return records
# ---------------------------------------------------------------------------
# DataFrame helpers (dashboard / analytics consumers)
# ---------------------------------------------------------------------------
_NUMERIC_COLS = ("latency", "throughput", "memory")
def normalize_dataframe(df: pd.DataFrame) -> pd.DataFrame:
"""Apply standard type coercions to a raw records DataFrame.
- Parses ``timestamp`` to UTC-aware datetime.
- Coerces ``latency``, ``throughput``, ``memory`` to float.
- Adds a ``config_id`` column (first 7 chars of ``commit_sha``).
Returns the mutated DataFrame (also modifies in place for efficiency).
"""
if df.empty:
return df
df["timestamp"] = pd.to_datetime(df["timestamp"], utc=True, errors="coerce")
df["config_id"] = df.get("commit_sha", pd.Series(dtype=str)).fillna("unknown").str[:7]
for col in _NUMERIC_COLS:
if col in df.columns:
df[col] = pd.to_numeric(df[col], errors="coerce")
return df
def load_as_dataframe(
local_dir: str,
*,
days: int | None = None,
successful_only: bool = False,
) -> pd.DataFrame:
"""Load and normalize records from *local_dir* into a pandas DataFrame.
Combines :func:`load_records` + :func:`normalize_dataframe` into a single
call for consumers (e.g. the dashboard) that work exclusively with
DataFrames.
Args:
local_dir: Root directory previously populated by :func:`sync_from_hf`.
days: Passed through to :func:`load_records`.
successful_only: Passed through to :func:`load_records`.
Returns:
Normalized DataFrame, or an empty DataFrame if no records were found.
"""
records = load_records(local_dir, days=days, successful_only=successful_only)
if not records:
return pd.DataFrame()
df = pd.DataFrame(records)
return normalize_dataframe(df)
@@ -164,10 +164,6 @@ def test_inference_performance(cfg):
avg_time = sum(times) / len(times)
max_peak_memory = max(peak_memories)
device_name = torch.cuda.get_device_name()
num_frames = gen_kwargs.get("num_frames")
throughput_fps = (1.0 / avg_time) if avg_time > 0 else None
if isinstance(num_frames, (int, float)) and avg_time > 0:
throughput_fps = num_frames / avg_time
results = {
"benchmark_id": cfg["benchmark_id"],
@@ -178,8 +174,6 @@ def test_inference_performance(cfg):
"num_measurement_runs": num_measure,
"avg_generation_time_s": round(avg_time, 3),
"individual_times_s": [round(t, 3) for t in times],
"throughput_fps": round(throughput_fps, 3)
if throughput_fps is not None else None,
"max_peak_memory_mb": round(max_peak_memory, 1),
"individual_peak_memories_mb": [round(m, 1) for m in peak_memories],
"thresholds": thresholds,
+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)
@@ -1 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
@@ -1,290 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.callback`.
Covers the ``Callback`` base class no-op contract and the
``CallbackDict`` instantiation / dispatch / state-dict logic.
The concrete callback subclasses (``GradNormClipCallback``,
``EMACallback``, ``ValidationCallback``) have their own test files.
"""
from __future__ import annotations
import logging
from collections.abc import Iterator
from contextlib import contextmanager
from typing import Any
import pytest
from fastvideo.train.callbacks.callback import (
Callback,
CallbackDict,
_BUILTIN_CALLBACKS,
)
# ``fastvideo.logger.init_logger`` sets ``propagate=False`` on its
# loggers, so the standard ``caplog`` fixture cannot observe them.
# This helper attaches a temporary handler directly to the target
# logger and yields the captured records.
@contextmanager
def _capture_logger(
name: str, level: int = logging.WARNING
) -> Iterator[list[logging.LogRecord]]:
logger = logging.getLogger(name)
records: list[logging.LogRecord] = []
class _Handler(logging.Handler):
def emit(self, record: logging.LogRecord) -> None:
records.append(record)
handler = _Handler(level=level)
prev_level = logger.level
logger.addHandler(handler)
logger.setLevel(level)
try:
yield records
finally:
logger.removeHandler(handler)
logger.setLevel(prev_level)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class _RecordingCallback(Callback):
"""Callback that records every hook call into a shared list."""
def __init__(self, *, tag: str, sink: list[str]) -> None:
self._tag = tag
self._sink = sink
def on_train_start(self, method, iteration: int = 0) -> None:
self._sink.append(f"{self._tag}:on_train_start:{iteration}")
def on_training_step_end(
self, method, loss_dict, iteration: int = 0
) -> None:
self._sink.append(f"{self._tag}:on_training_step_end:{iteration}")
def on_validation_begin(self, method, iteration: int = 0) -> None:
self._sink.append(f"{self._tag}:on_validation_begin:{iteration}")
def state_dict(self) -> dict[str, Any]:
return {"tag": self._tag, "marker": 7}
def load_state_dict(self, sd: dict[str, Any]) -> None:
self._sink.append(f"{self._tag}:load:{sd.get('marker')}")
class _NotACallback:
"""Plain class used to exercise the non-Callback type guard."""
def __init__(self, **_: Any) -> None:
pass
# ---------------------------------------------------------------------------
# A. Callback base class
# ---------------------------------------------------------------------------
class TestCallbackBase:
def test_default_hooks_return_none(self) -> None:
cb = Callback()
assert cb.on_train_start(method=None) is None
assert (
cb.on_training_step_end(method=None, loss_dict={}) is None
)
assert cb.on_before_optimizer_step(method=None) is None
assert cb.on_validation_begin(method=None) is None
assert cb.on_validation_end(method=None) is None
assert cb.on_train_end(method=None) is None
def test_default_state_dict_round_trip(self) -> None:
cb = Callback()
assert cb.state_dict() == {}
# Default load_state_dict accepts arbitrary state without raising.
assert cb.load_state_dict({"unrelated": 1}) is None
# ---------------------------------------------------------------------------
# B. CallbackDict construction
# ---------------------------------------------------------------------------
class TestCallbackDictInit:
def test_empty_config(self) -> None:
cb_dict = CallbackDict({}, training_config=object())
assert cb_dict._callbacks == {}
def test_builtin_name_resolves_without_target(self) -> None:
# ``grad_clip`` is a registered builtin.
cfg = {"grad_clip": {"max_grad_norm": 0.5}}
tc = object()
cb_dict = CallbackDict(cfg, training_config=tc)
assert "grad_clip" in cb_dict._callbacks
from fastvideo.train.callbacks.grad_clip import (
GradNormClipCallback,
)
cb = cb_dict._callbacks["grad_clip"]
assert isinstance(cb, GradNormClipCallback)
# CallbackDict wires up training_config + back-pointer.
assert cb.training_config is tc
assert cb._callback_dict is cb_dict
def test_explicit_target_overrides_name_lookup(self) -> None:
cfg = {
"anything_goes": {
"_target_": (
"fastvideo.train.callbacks.grad_clip."
"GradNormClipCallback"
),
"max_grad_norm": 1.0,
}
}
cb_dict = CallbackDict(cfg, training_config=object())
assert "anything_goes" in cb_dict._callbacks
def test_unknown_name_without_target_is_skipped(self) -> None:
cfg = {"mystery": {"some_arg": 1}}
with _capture_logger(
"fastvideo.train.callbacks.callback"
) as records:
cb_dict = CallbackDict(cfg, training_config=object())
assert cb_dict._callbacks == {}
assert any(
"missing" in r.getMessage() and "mystery" in r.getMessage()
for r in records
)
def test_non_callback_target_raises(self) -> None:
cfg = {
"bad": {
"_target_": (
"fastvideo.tests.train.callbacks.test_callback."
"_NotACallback"
)
}
}
with pytest.raises(TypeError, match="expected a Callback"):
CallbackDict(cfg, training_config=object())
def test_builtin_registry_has_expected_entries(self) -> None:
# Sanity: protect the builtin registry from silent shrinkage.
assert set(_BUILTIN_CALLBACKS) >= {
"grad_clip",
"validation",
"ema",
}
# ---------------------------------------------------------------------------
# C. Dispatch via __getattr__
# ---------------------------------------------------------------------------
class TestCallbackDictDispatch:
def _build(
self, sink: list[str]
) -> CallbackDict:
cb_dict = CallbackDict({}, training_config=object())
cb_dict._callbacks["first"] = _RecordingCallback(
tag="first", sink=sink
)
cb_dict._callbacks["second"] = _RecordingCallback(
tag="second", sink=sink
)
return cb_dict
def test_dispatch_calls_all_in_insertion_order(self) -> None:
sink: list[str] = []
cb_dict = self._build(sink)
cb_dict.on_train_start(method=None, iteration=3)
assert sink == [
"first:on_train_start:3",
"second:on_train_start:3",
]
def test_dispatch_to_hook_some_callbacks_skip(self) -> None:
sink: list[str] = []
cb_dict = self._build(sink)
# The base Callback subclass below only implements one hook;
# dispatch should still fan out without raising.
class _OnlyValidation(Callback):
def on_validation_end(
self, method, iteration: int = 0
) -> None:
sink.append(f"vend:{iteration}")
cb_dict._callbacks["only_v"] = _OnlyValidation()
cb_dict.on_validation_end(method=None, iteration=11)
assert "vend:11" in sink
def test_dispatch_unknown_hook_is_noop(self) -> None:
# Methods that don't exist on any callback should not raise.
cb_dict = self._build([])
cb_dict.totally_made_up_hook(method=None, iteration=0)
def test_underscore_attribute_raises(self) -> None:
cb_dict = self._build([])
with pytest.raises(AttributeError):
getattr(cb_dict, "_does_not_exist")
# ---------------------------------------------------------------------------
# D. state_dict / load_state_dict
# ---------------------------------------------------------------------------
class TestCallbackDictStateDict:
def _build(self) -> tuple[CallbackDict, list[str]]:
sink: list[str] = []
cb_dict = CallbackDict({}, training_config=object())
cb_dict._callbacks["first"] = _RecordingCallback(
tag="first", sink=sink
)
cb_dict._callbacks["second"] = _RecordingCallback(
tag="second", sink=sink
)
return cb_dict, sink
def test_state_dict_returns_per_callback_dict(self) -> None:
cb_dict, _ = self._build()
state = cb_dict.state_dict()
assert set(state) == {"first", "second"}
assert state["first"] == {"tag": "first", "marker": 7}
assert state["second"] == {"tag": "second", "marker": 7}
def test_load_state_dict_dispatches_to_each(self) -> None:
cb_dict, sink = self._build()
cb_dict.load_state_dict(
{
"first": {"marker": 1},
"second": {"marker": 2},
}
)
assert sink == ["first:load:1", "second:load:2"]
def test_load_state_dict_missing_key_warns_no_raise(self) -> None:
cb_dict, sink = self._build()
with _capture_logger(
"fastvideo.train.callbacks.callback"
) as records:
cb_dict.load_state_dict({"first": {"marker": 99}})
assert sink == ["first:load:99"]
assert any(
"second" in r.getMessage() and "not found" in r.getMessage()
for r in records
)
-254
View File
@@ -1,254 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.ema`.
Exercises the EMA lifecycle (lazy init, ``start_iter`` gating, decay
math, ``ema_context`` swap, state-dict round-trip) on a tiny CPU
``nn.Linear``. ``EMA_FSDP`` works without ``dist.init_process_group``
because ``dist.is_initialized()`` returns False and ``_to_local_tensor``
falls through to raw tensors for non-DTensor inputs.
"""
from __future__ import annotations
from typing import Any
import pytest
import torch
from fastvideo.train.callbacks.ema import EMACallback
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class _Student:
def __init__(self, transformer: torch.nn.Module | None) -> None:
self.transformer = transformer
class _RecordingTracker:
def __init__(self) -> None:
self.entries: list[tuple[dict[str, Any], int]] = []
def log(self, payload: dict[str, Any], step: int) -> None:
self.entries.append((payload, step))
class _Method:
def __init__(
self,
transformer: torch.nn.Module | None,
tracker: Any | None = None,
) -> None:
self.student = _Student(transformer)
self.tracker = tracker
def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module:
m = torch.nn.Linear(4, 2, bias=False)
with torch.no_grad():
m.weight.fill_(fill)
return m
# ---------------------------------------------------------------------------
# A. on_train_start
# ---------------------------------------------------------------------------
class TestOnTrainStart:
def test_initializes_ema_from_student(self) -> None:
transformer = _tiny_transformer(fill=0.5)
cb = EMACallback(decay=0.9, start_iter=0)
cb.on_train_start(_Method(transformer), iteration=0)
assert cb.student_ema is not None
# Shadow shape matches transformer parameter.
shadow = cb.student_ema.shadow["weight"]
assert shadow.shape == transformer.weight.shape
assert torch.allclose(shadow, transformer.weight.detach().cpu())
def test_missing_transformer_raises(self) -> None:
cb = EMACallback()
with pytest.raises(ValueError, match="No student transformer"):
cb.on_train_start(_Method(transformer=None), iteration=0)
# ---------------------------------------------------------------------------
# B. on_training_step_end (decay math + start_iter gating)
# ---------------------------------------------------------------------------
class TestOnTrainingStepEnd:
def test_no_op_before_train_start(self) -> None:
cb = EMACallback()
# student_ema is None until on_train_start.
cb.on_training_step_end(
_Method(transformer=None), loss_dict={}, iteration=0
)
assert not cb._ema_started
def test_skipped_until_start_iter(self) -> None:
transformer = _tiny_transformer(fill=1.0)
cb = EMACallback(decay=0.5, start_iter=10)
cb.on_train_start(_Method(transformer), iteration=0)
# Mutate transformer to drift it away from initial shadow.
with torch.no_grad():
transformer.weight.fill_(7.0)
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=5
)
# Below start_iter: shadow is untouched, _ema_started False.
assert not cb._ema_started
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 1.0),
)
def test_first_active_step_reinits_then_updates(self) -> None:
transformer = _tiny_transformer(fill=1.0)
cb = EMACallback(decay=0.9, start_iter=10)
cb.on_train_start(_Method(transformer), iteration=0)
# Drift transformer so that re-init has a visible effect.
with torch.no_grad():
transformer.weight.fill_(5.0)
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=10
)
# First active step: shadow is re-initialized from the
# current transformer (5.0) and *then* update() applies decay
# against the same value, so shadow stays at 5.0.
assert cb._ema_started
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 5.0),
)
def test_subsequent_step_applies_decay(self) -> None:
transformer = _tiny_transformer(fill=2.0)
cb = EMACallback(decay=0.9, start_iter=0)
cb.on_train_start(_Method(transformer), iteration=0)
# Step 0: re-init at 2.0, then update against 2.0 → still 2.0.
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=0
)
# Step 1: drift transformer to 12.0, expect
# shadow = 0.9 * 2.0 + 0.1 * 12.0 = 3.0.
with torch.no_grad():
transformer.weight.fill_(12.0)
cb.on_training_step_end(
_Method(transformer), loss_dict={}, iteration=1
)
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 3.0),
atol=1e-6,
)
def test_tracker_logs_decay(self) -> None:
transformer = _tiny_transformer()
tracker = _RecordingTracker()
cb = EMACallback(decay=0.99, start_iter=0)
method = _Method(transformer, tracker=tracker)
cb.on_train_start(method, iteration=0)
cb.on_training_step_end(method, loss_dict={}, iteration=0)
assert any(
payload.get("ema/decay") == 0.99 and step == 0
for payload, step in tracker.entries
)
# ---------------------------------------------------------------------------
# C. ema_context
# ---------------------------------------------------------------------------
class TestEmaContext:
def test_passthrough_when_inactive(self) -> None:
transformer = _tiny_transformer(fill=3.0)
cb = EMACallback()
# No on_train_start → student_ema is None.
with cb.ema_context(transformer) as t:
assert t is transformer
assert torch.allclose(
t.weight, torch.full((2, 4), 3.0)
)
def test_swaps_weights_then_restores(self) -> None:
transformer = _tiny_transformer(fill=1.0)
cb = EMACallback(decay=0.0, start_iter=0)
method = _Method(transformer)
cb.on_train_start(method, iteration=0)
# decay=0 → after one step the shadow == current weights == 1.0.
cb.on_training_step_end(method, loss_dict={}, iteration=0)
# Drift transformer; ema_context should swap shadow (1.0) in
# for the duration and restore the post-drift value (9.0).
with torch.no_grad():
transformer.weight.fill_(9.0)
with cb.ema_context(transformer) as t:
assert torch.allclose(t.weight, torch.full((2, 4), 1.0))
assert torch.allclose(
transformer.weight, torch.full((2, 4), 9.0)
)
# ---------------------------------------------------------------------------
# D. State dict round-trip
# ---------------------------------------------------------------------------
class TestStateDict:
def test_state_dict_empty_before_train_start(self) -> None:
cb = EMACallback()
assert cb.state_dict() == {}
def test_round_trip_preserves_shadow_and_started_flag(self) -> None:
transformer = _tiny_transformer(fill=4.0)
cb = EMACallback(decay=0.5, start_iter=0)
method = _Method(transformer)
cb.on_train_start(method, iteration=0)
cb.on_training_step_end(method, loss_dict={}, iteration=0)
state = cb.state_dict()
assert "student_ema" in state
assert state["ema_started"] is True
# Build a fresh callback and load.
fresh = EMACallback(decay=0.5, start_iter=0)
fresh.on_train_start(_Method(_tiny_transformer(fill=0.0)),
iteration=0)
# Sanity: fresh shadow != saved shadow before load.
assert not torch.allclose(
fresh.student_ema.shadow["weight"],
cb.student_ema.shadow["weight"],
)
fresh.load_state_dict(state)
assert fresh._ema_started is True
assert torch.allclose(
fresh.student_ema.shadow["weight"],
cb.student_ema.shadow["weight"],
)
def test_load_without_student_ema_only_sets_flag(self) -> None:
cb = EMACallback()
# student_ema is None — load must not attempt to assign shadow.
cb.load_state_dict({"ema_started": True})
assert cb._ema_started is True
assert cb.student_ema is None
@@ -1,170 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.grad_clip`.
Exercises ``GradNormClipCallback.on_before_optimizer_step`` against
synthetic ``nn.Module`` targets with manually populated gradients.
"""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class _RecordingTracker:
"""Tracker stub that records every ``log`` call."""
def __init__(self) -> None:
self.entries: list[tuple[dict[str, Any], int]] = []
def log(self, payload: dict[str, Any], step: int) -> None:
self.entries.append((payload, step))
class _Method:
"""Minimal stand-in for ``TrainingMethod``."""
def __init__(
self,
targets: dict[str, torch.nn.Module],
tracker: Any | None = None,
) -> None:
self._targets = targets
self.tracker = tracker
self.iter_seen: int | None = None
def get_grad_clip_targets(
self, iteration: int
) -> dict[str, torch.nn.Module]:
self.iter_seen = iteration
return self._targets
def _make_module(*, grad_value: float, n: int = 4) -> torch.nn.Module:
"""Return an ``nn.Linear`` whose grads are filled with ``grad_value``."""
m = torch.nn.Linear(n, n, bias=False)
m.weight.grad = torch.full_like(m.weight, fill_value=grad_value)
return m
def _grad_norm(module: torch.nn.Module) -> float:
flat = torch.cat([p.grad.flatten() for p in module.parameters()])
return float(flat.norm(2).item())
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestGradNormClipCallback:
def test_disabled_when_max_norm_non_positive(self) -> None:
m = _make_module(grad_value=10.0)
before = _grad_norm(m)
cb = GradNormClipCallback(max_grad_norm=0.0)
method = _Method(targets={"m": m})
cb.on_before_optimizer_step(method=method, iteration=0)
# No clipping applied; ``get_grad_clip_targets`` not consulted.
assert _grad_norm(m) == before
assert method.iter_seen is None
def test_large_grads_get_clipped(self) -> None:
m = _make_module(grad_value=10.0)
assert _grad_norm(m) > 1.0
cb = GradNormClipCallback(max_grad_norm=1.0)
cb.on_before_optimizer_step(
method=_Method(targets={"m": m}), iteration=0
)
# After clipping the L2 norm should not exceed max_grad_norm
# (allow a tiny epsilon for the +1e-6 in the clip helper).
assert _grad_norm(m) <= 1.0 + 1e-4
def test_small_grads_unchanged(self) -> None:
m = _make_module(grad_value=0.01)
before = _grad_norm(m)
assert before < 1.0
cb = GradNormClipCallback(max_grad_norm=1.0)
cb.on_before_optimizer_step(
method=_Method(targets={"m": m}), iteration=0
)
# Clip coef >1 is clamped to 1, so values are preserved
# (modulo the *1.0 multiply, which is exact for floats).
assert abs(_grad_norm(m) - before) < 1e-6
def test_iteration_forwarded_to_targets(self) -> None:
m = _make_module(grad_value=0.5)
cb = GradNormClipCallback(max_grad_norm=1.0)
method = _Method(targets={"m": m})
cb.on_before_optimizer_step(method=method, iteration=42)
assert method.iter_seen == 42
def test_tracker_logged_when_enabled(self) -> None:
m = _make_module(grad_value=5.0)
tracker = _RecordingTracker()
cb = GradNormClipCallback(
max_grad_norm=1.0, log_grad_norms=True
)
cb.on_before_optimizer_step(
method=_Method(targets={"layer": m}, tracker=tracker),
iteration=7,
)
assert len(tracker.entries) == 1
payload, step = tracker.entries[0]
assert step == 7
assert "grad_norm/layer" in payload
assert payload["grad_norm/layer"] > 0.0
def test_tracker_not_logged_when_disabled(self) -> None:
m = _make_module(grad_value=5.0)
tracker = _RecordingTracker()
cb = GradNormClipCallback(
max_grad_norm=1.0, log_grad_norms=False
)
cb.on_before_optimizer_step(
method=_Method(targets={"m": m}, tracker=tracker),
iteration=0,
)
assert tracker.entries == []
def test_no_tracker_does_not_raise(self) -> None:
m = _make_module(grad_value=5.0)
cb = GradNormClipCallback(max_grad_norm=1.0, log_grad_norms=True)
# Method without a tracker attribute at all.
class _BareMethod:
def get_grad_clip_targets(
self, iteration: int
) -> dict[str, torch.nn.Module]:
return {"m": m}
cb.on_before_optimizer_step(method=_BareMethod(), iteration=0)
# No assertion — must simply not raise.
def test_multiple_targets_each_logged(self) -> None:
targets = {
"head": _make_module(grad_value=4.0),
"tail": _make_module(grad_value=8.0),
}
tracker = _RecordingTracker()
cb = GradNormClipCallback(max_grad_norm=1.0)
cb.on_before_optimizer_step(
method=_Method(targets=targets, tracker=tracker),
iteration=1,
)
keys = {next(iter(p)) for p, _ in tracker.entries}
assert keys == {"grad_norm/head", "grad_norm/tail"}
@@ -1,234 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.validation`.
Covers the parts of ``ValidationCallback`` that don't need a real
pipeline or distributed init:
* constructor type coercions and defaults,
* ``on_validation_begin`` gating logic (every_steps + modulo),
* ``_find_ema_callback`` lookup via ``_callback_dict``,
* ``state_dict`` / ``load_state_dict`` rng round-trip.
The heavy ``_run_validation`` path needs a real diffusion pipeline plus
distributed init and is exercised by Phase 2/3 tests.
"""
from __future__ import annotations
import torch
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.callbacks.ema import EMACallback
from fastvideo.train.callbacks.validation import ValidationCallback
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
_PIPE_TARGET = "fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline"
def _make_callback(
*,
every_steps: int = 100,
sampling_steps: list[int] | None = None,
guidance_scale: float | None = None,
num_frames: int | None = None,
sampling_timesteps: list[int] | None = None,
output_dir: str | None = None,
) -> ValidationCallback:
return ValidationCallback(
pipeline_target=_PIPE_TARGET,
dataset_file="/tmp/does_not_exist.json",
every_steps=every_steps,
sampling_steps=sampling_steps,
guidance_scale=guidance_scale,
num_frames=num_frames,
sampling_timesteps=sampling_timesteps,
output_dir=output_dir,
)
# ---------------------------------------------------------------------------
# A. Constructor coercions / defaults
# ---------------------------------------------------------------------------
class TestConstructor:
def test_defaults(self) -> None:
cb = _make_callback()
assert cb.pipeline_target == _PIPE_TARGET
assert cb.dataset_file == "/tmp/does_not_exist.json"
assert cb.every_steps == 100
assert cb.sampling_steps == [40]
assert cb.guidance_scale is None
assert cb.num_frames is None
assert cb.sampling_timesteps is None
assert cb.output_dir is None
# Lazy fields not yet populated.
assert cb._pipeline is None
assert cb._sampling_param is None
assert cb.validation_random_generator is None
def test_string_inputs_are_coerced(self) -> None:
# YAML often produces strings for numeric fields; the
# constructor must coerce them.
cb = ValidationCallback(
pipeline_target=_PIPE_TARGET,
dataset_file="x.json",
every_steps="50", # type: ignore[arg-type]
sampling_steps=["20", "40"], # type: ignore[arg-type]
guidance_scale="4.5", # type: ignore[arg-type]
num_frames="77", # type: ignore[arg-type]
sampling_timesteps=["1000", "500"],
)
assert cb.every_steps == 50
assert cb.sampling_steps == [20, 40]
assert cb.guidance_scale == 4.5
assert cb.num_frames == 77
assert cb.sampling_timesteps == [1000, 500]
def test_pipeline_kwargs_collected(self) -> None:
cb = ValidationCallback(
pipeline_target=_PIPE_TARGET,
dataset_file="x.json",
extra_arg=123,
another="value",
)
# Unknown kwargs are stashed for the pipeline factory.
assert cb.pipeline_kwargs == {
"extra_arg": 123,
"another": "value",
}
# ---------------------------------------------------------------------------
# B. on_validation_begin gating
# ---------------------------------------------------------------------------
class _NoRunValidation(ValidationCallback):
"""Subclass that records ``_run_validation`` calls instead of
actually running them — lets us assert the gating logic without a
real pipeline."""
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.run_calls: list[int] = []
def _run_validation(self, method, step: int) -> None: # type: ignore[override]
self.run_calls.append(step)
def _make_recording(**kwargs) -> _NoRunValidation:
return _NoRunValidation(
pipeline_target=_PIPE_TARGET,
dataset_file="x.json",
**kwargs,
)
class TestOnValidationBegin:
def test_skipped_when_every_steps_zero(self) -> None:
cb = _make_recording(every_steps=0)
cb.on_validation_begin(method=None, iteration=0)
cb.on_validation_begin(method=None, iteration=1000)
assert cb.run_calls == []
def test_skipped_on_off_iter(self) -> None:
cb = _make_recording(every_steps=50)
cb.on_validation_begin(method=None, iteration=49)
cb.on_validation_begin(method=None, iteration=51)
assert cb.run_calls == []
def test_runs_on_match(self) -> None:
cb = _make_recording(every_steps=50)
cb.on_validation_begin(method=None, iteration=50)
cb.on_validation_begin(method=None, iteration=100)
assert cb.run_calls == [50, 100]
def test_iter_zero_runs(self) -> None:
# 0 % anything == 0 → step 0 fires (matches existing
# validation behavior used by ValidationCallback consumers).
cb = _make_recording(every_steps=50)
cb.on_validation_begin(method=None, iteration=0)
assert cb.run_calls == [0]
# ---------------------------------------------------------------------------
# C. _find_ema_callback
# ---------------------------------------------------------------------------
class TestFindEmaCallback:
def test_returns_none_without_callback_dict(self) -> None:
cb = _make_callback()
# _callback_dict is not set on bare instances.
assert cb._find_ema_callback() is None
def test_returns_none_when_no_ema_registered(self) -> None:
cb = _make_callback()
cb_dict = CallbackDict({}, training_config=object())
cb._callback_dict = cb_dict
assert cb._find_ema_callback() is None
def test_finds_ema_callback(self) -> None:
cb = _make_callback()
cb_dict = CallbackDict({}, training_config=object())
ema = EMACallback(decay=0.99)
cb_dict._callbacks["ema"] = ema
cb_dict._callbacks["validation"] = cb
cb._callback_dict = cb_dict
found = cb._find_ema_callback()
assert found is ema
# ---------------------------------------------------------------------------
# D. state_dict / load_state_dict (rng round-trip)
# ---------------------------------------------------------------------------
class TestStateDict:
def test_state_dict_empty_without_generator(self) -> None:
cb = _make_callback()
# validation_random_generator is None until on_train_start.
assert cb.state_dict() == {}
def test_round_trip_preserves_rng_state(self) -> None:
cb = _make_callback()
gen = torch.Generator(device="cpu").manual_seed(123)
# Advance RNG so a default-init generator on the receiving
# side is observably different.
for _ in range(5):
torch.randn(4, generator=gen)
cb.validation_random_generator = gen
state = cb.state_dict()
assert "validation_rng" in state
# Receiver: fresh generator with a different seed.
fresh = _make_callback()
fresh.validation_random_generator = (
torch.Generator(device="cpu").manual_seed(999)
)
fresh.load_state_dict(state)
# After load, both generators draw the same next sample.
a = torch.randn(8, generator=cb.validation_random_generator)
b = torch.randn(8, generator=fresh.validation_random_generator)
assert torch.equal(a, b)
def test_load_without_generator_is_noop(self) -> None:
cb = _make_callback()
# Generator is None: load must not raise even when state has
# an rng entry.
cb.load_state_dict(
{"validation_rng": torch.tensor([1, 2, 3], dtype=torch.uint8)}
)
assert cb.validation_random_generator is None
@@ -1,360 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only unit tests for :mod:`fastvideo.train.utils.checkpoint`.
Covers the pure-Python portions of the checkpoint manager: name
parsing, resume-path resolution, metadata round-trip, rolling-delete
cleanup, the ``_is_stateful`` predicate, and the ``maybe_save`` gating
logic. Code paths that touch DCP (``dcp.save`` / ``dcp.load``) and
CUDA RNG snapshots are intentionally not covered here — those need a
GPU runner and will be tested in later phases.
"""
from __future__ import annotations
from pathlib import Path
from typing import Any
import pytest
from fastvideo.train.utils.checkpoint import (
CheckpointConfig,
CheckpointManager,
_find_latest_checkpoint,
_is_stateful,
_parse_step_from_dir,
_resolve_resume_checkpoint,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_checkpoint_dir(
output_dir: Path,
step: int,
*,
with_dcp: bool = True,
) -> Path:
"""Create a fake ``checkpoint-<step>/dcp`` directory tree."""
ckpt_dir = output_dir / f"checkpoint-{step}"
ckpt_dir.mkdir(parents=True, exist_ok=True)
if with_dcp:
(ckpt_dir / "dcp").mkdir(exist_ok=True)
return ckpt_dir
def _make_manager(
tmp_path: Path,
*,
save_steps: int = 0,
keep_last: int = 0,
raw_config: dict[str, Any] | None = None,
) -> CheckpointManager:
"""Build a minimal ``CheckpointManager`` for tests that don't touch DCP."""
return CheckpointManager(
method=None,
dataloader=None,
output_dir=str(tmp_path),
config=CheckpointConfig(save_steps=save_steps, keep_last=keep_last),
raw_config=raw_config,
)
# ---------------------------------------------------------------------------
# A. _is_stateful predicate
# ---------------------------------------------------------------------------
class _Full:
def state_dict(self) -> dict[str, Any]:
return {}
def load_state_dict(self, sd: dict[str, Any]) -> None:
pass
class _MissingStateDict:
def load_state_dict(self, sd: dict[str, Any]) -> None:
pass
class _MissingLoad:
def state_dict(self) -> dict[str, Any]:
return {}
def test_is_stateful_true_for_full_object() -> None:
assert _is_stateful(_Full()) is True
def test_is_stateful_false_when_missing_state_dict() -> None:
assert _is_stateful(_MissingStateDict()) is False
def test_is_stateful_false_when_missing_load_state_dict() -> None:
assert _is_stateful(_MissingLoad()) is False
# ---------------------------------------------------------------------------
# B. _parse_step_from_dir
# ---------------------------------------------------------------------------
def test_parse_step_valid(tmp_path: Path) -> None:
assert _parse_step_from_dir(tmp_path / "checkpoint-100") == 100
def test_parse_step_zero(tmp_path: Path) -> None:
assert _parse_step_from_dir(tmp_path / "checkpoint-0") == 0
def test_parse_step_invalid_raises(tmp_path: Path) -> None:
with pytest.raises(ValueError, match="Invalid checkpoint directory"):
_parse_step_from_dir(tmp_path / "not-a-checkpoint")
# ---------------------------------------------------------------------------
# C. _find_latest_checkpoint
# ---------------------------------------------------------------------------
def test_find_latest_returns_none_on_nonexistent_dir(tmp_path: Path) -> None:
assert _find_latest_checkpoint(tmp_path / "missing") is None
def test_find_latest_returns_none_on_empty_dir(tmp_path: Path) -> None:
assert _find_latest_checkpoint(tmp_path) is None
def test_find_latest_returns_largest_step(tmp_path: Path) -> None:
_make_checkpoint_dir(tmp_path, 10)
_make_checkpoint_dir(tmp_path, 200)
_make_checkpoint_dir(tmp_path, 50)
latest = _find_latest_checkpoint(tmp_path)
assert latest is not None
assert latest.name == "checkpoint-200"
def test_find_latest_skips_dirs_without_dcp_subdir(tmp_path: Path) -> None:
# checkpoint-10 is "corrupted" — has no dcp/ subdir, must be skipped.
_make_checkpoint_dir(tmp_path, 10, with_dcp=False)
_make_checkpoint_dir(tmp_path, 5, with_dcp=True)
latest = _find_latest_checkpoint(tmp_path)
assert latest is not None
assert latest.name == "checkpoint-5"
def test_find_latest_skips_non_checkpoint_dirs(tmp_path: Path) -> None:
(tmp_path / "logs").mkdir()
(tmp_path / "wandb").mkdir()
(tmp_path / "some_file.txt").write_text("noise")
_make_checkpoint_dir(tmp_path, 7)
latest = _find_latest_checkpoint(tmp_path)
assert latest is not None
assert latest.name == "checkpoint-7"
# ---------------------------------------------------------------------------
# D. _resolve_resume_checkpoint
# ---------------------------------------------------------------------------
def test_resolve_latest_with_no_checkpoints_returns_none(
tmp_path: Path) -> None:
out = tmp_path / "outputs"
out.mkdir()
assert _resolve_resume_checkpoint("latest", output_dir=str(out)) is None
def test_resolve_latest_returns_latest_checkpoint(tmp_path: Path) -> None:
_make_checkpoint_dir(tmp_path, 30)
_make_checkpoint_dir(tmp_path, 10)
resolved = _resolve_resume_checkpoint("latest", output_dir=str(tmp_path))
assert resolved is not None
assert resolved.name == "checkpoint-30"
def test_resolve_explicit_checkpoint_dir(tmp_path: Path) -> None:
ckpt = _make_checkpoint_dir(tmp_path, 42)
resolved = _resolve_resume_checkpoint(str(ckpt),
output_dir=str(tmp_path))
assert resolved is not None
assert resolved.name == "checkpoint-42"
def test_resolve_dcp_subdir_returns_parent_checkpoint(tmp_path: Path) -> None:
ckpt = _make_checkpoint_dir(tmp_path, 42)
dcp_path = ckpt / "dcp"
resolved = _resolve_resume_checkpoint(str(dcp_path),
output_dir=str(tmp_path))
assert resolved is not None
assert resolved.name == "checkpoint-42"
def test_resolve_output_dir_returns_latest(tmp_path: Path) -> None:
out = tmp_path / "outputs"
out.mkdir()
_make_checkpoint_dir(out, 100)
_make_checkpoint_dir(out, 50)
resolved = _resolve_resume_checkpoint(str(out), output_dir=str(tmp_path))
assert resolved is not None
assert resolved.name == "checkpoint-100"
def test_resolve_nonexistent_path_raises(tmp_path: Path) -> None:
with pytest.raises(FileNotFoundError):
_resolve_resume_checkpoint(str(tmp_path / "missing"),
output_dir=str(tmp_path))
def test_resolve_checkpoint_without_dcp_raises(tmp_path: Path) -> None:
ckpt = _make_checkpoint_dir(tmp_path, 5, with_dcp=False)
with pytest.raises(FileNotFoundError, match="dcp"):
_resolve_resume_checkpoint(str(ckpt), output_dir=str(tmp_path))
def test_resolve_unknown_dir_raises(tmp_path: Path) -> None:
"""A dir that is neither a checkpoint nor an output_dir-with-checkpoints."""
bogus = tmp_path / "bogus"
bogus.mkdir()
with pytest.raises(ValueError, match="Could not resolve"):
_resolve_resume_checkpoint(str(bogus), output_dir=str(tmp_path))
# ---------------------------------------------------------------------------
# E. metadata read/write
# ---------------------------------------------------------------------------
def test_write_metadata_roundtrip_with_step(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path)
ckpt_dir = _make_checkpoint_dir(tmp_path, 7)
mgr._write_metadata(ckpt_dir, step=7)
loaded = CheckpointManager.load_metadata(ckpt_dir)
assert loaded == {"step": 7}
def test_write_metadata_includes_raw_config(tmp_path: Path) -> None:
raw = {
"models": {
"student": {
"_target_": "X"
}
},
"training": {
"distributed": {
"num_gpus": 4
}
},
}
mgr = _make_manager(tmp_path, raw_config=raw)
ckpt_dir = _make_checkpoint_dir(tmp_path, 7)
mgr._write_metadata(ckpt_dir, step=7)
loaded = CheckpointManager.load_metadata(ckpt_dir)
assert loaded["step"] == 7
assert loaded["config"] == raw
def test_load_metadata_raises_on_missing_file(tmp_path: Path) -> None:
ckpt_dir = _make_checkpoint_dir(tmp_path, 7)
# No metadata.json written.
with pytest.raises(FileNotFoundError, match="metadata"):
CheckpointManager.load_metadata(ckpt_dir)
# ---------------------------------------------------------------------------
# F. _cleanup_old_checkpoints (rolling delete)
# ---------------------------------------------------------------------------
def test_cleanup_keep_last_zero_is_noop(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, keep_last=0)
for step in (1, 2, 3):
_make_checkpoint_dir(tmp_path, step)
mgr._cleanup_old_checkpoints()
remaining = sorted(p.name for p in tmp_path.iterdir())
assert remaining == ["checkpoint-1", "checkpoint-2", "checkpoint-3"]
def test_cleanup_keeps_newest_when_over_limit(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, keep_last=2)
for step in (1, 5, 10, 50, 100):
_make_checkpoint_dir(tmp_path, step)
mgr._cleanup_old_checkpoints()
remaining = sorted(p.name for p in tmp_path.iterdir())
assert remaining == ["checkpoint-100", "checkpoint-50"]
def test_cleanup_no_op_when_under_limit(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, keep_last=3)
for step in (1, 2):
_make_checkpoint_dir(tmp_path, step)
mgr._cleanup_old_checkpoints()
remaining = sorted(p.name for p in tmp_path.iterdir())
assert remaining == ["checkpoint-1", "checkpoint-2"]
def test_cleanup_skips_non_checkpoint_dirs(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, keep_last=1)
for step in (1, 2, 3):
_make_checkpoint_dir(tmp_path, step)
(tmp_path / "logs").mkdir()
(tmp_path / "wandb").mkdir()
mgr._cleanup_old_checkpoints()
remaining = sorted(p.name for p in tmp_path.iterdir())
assert remaining == ["checkpoint-3", "logs", "wandb"]
# ---------------------------------------------------------------------------
# G. maybe_save gating logic
# ---------------------------------------------------------------------------
def _record_save_calls(mgr: CheckpointManager) -> list[int]:
"""Replace ``mgr.save`` with a recorder that mimics the side effect
of advancing ``_last_saved_step`` so dedup logic still works."""
calls: list[int] = []
def fake_save(step: int) -> None:
calls.append(step)
mgr._last_saved_step = step
mgr.save = fake_save # type: ignore[method-assign]
return calls
def test_maybe_save_skipped_when_save_steps_is_zero(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, save_steps=0)
calls = _record_save_calls(mgr)
mgr.maybe_save(step=10)
mgr.maybe_save(step=100)
assert calls == []
def test_maybe_save_skipped_when_step_not_on_interval(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, save_steps=10)
calls = _record_save_calls(mgr)
for step in (1, 5, 9, 11, 15):
mgr.maybe_save(step=step)
assert calls == []
def test_maybe_save_dedupes_on_repeated_call(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, save_steps=10)
calls = _record_save_calls(mgr)
mgr.maybe_save(step=20)
mgr.maybe_save(step=20)
mgr.maybe_save(step=20)
assert calls == [20]
def test_maybe_save_triggers_on_each_interval(tmp_path: Path) -> None:
mgr = _make_manager(tmp_path, save_steps=10)
calls = _record_save_calls(mgr)
for step in range(1, 41):
mgr.maybe_save(step=step)
assert calls == [10, 20, 30, 40]

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