Compare commits

...
Author SHA1 Message Date
Satyam Srivastava b3a9874fc8 Merge branch 'main' into test-hf-sync 2026-05-01 11:38:11 -07:00
Satyam Srivastava d657cbbf17 Update regression gatekeep to 5% 2026-05-01 11:05:33 -07:00
William Lin 9801037c3d [refactor] tests/local_tests: organize by model family (#1269) 2026-05-01 01:49:54 -07:00
alexzmsandmergify[bot] 74d09b0efd [misc] cleanup: grad-norm asserts, dead offload file, callback names (#1268)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-05-01 01:16:13 -07:00
alexzms 38dc8820ac [ci] add CPU unit tests for train callback system in fastvideo.train (#1267) 2026-05-01 00:53:29 -07:00
William Lin c77a76c6af [feat] Stable Audio Open 1.0: T2A + A2A + RePaint inpainting (native) (#1260) 2026-05-01 00:07:11 -07:00
alexzms d14d5aadea [feat] Cosmos 2.5 training support in fastvideo.train (#1224) 2026-05-01 01:15:02 +00:00
alexzms 4c915b7742 [ci] add CPU unit tests for train checkpoint utilities in fastvideo.train (#1265) 2026-04-29 18:55:39 +00:00
Satyam Srivastava 48534ef4de [ci] Use median instead of mean to detect regressions. 2026-04-28 01:19:19 -07:00
Satyam Srivastava 1116f514be Refactor upload performance metrics 2026-04-28 01:07:27 -07:00
Satyam Srivastava d451e61749 Abstract hf code and fix plot bug 2026-04-27 23:12:42 -07:00
alexzms 9a8bbe18fa [bugfix]: fix SP deadlock in negative prompt encoding during training (#1178) 2026-04-28 01:06:49 +00:00
alexzms ea25441ef0 [ci] add CPU unit tests for fastvideo.train load_run_config (#1264) 2026-04-28 01:06:18 +00:00
48957fcde1 [bugfix] Fix modal remote functions crash container on sys exit in CI remote functions (#1261)
Co-authored-by: Satyam Srivastava <satyam53@Mac.lan1>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-27 21:50:55 +00:00
Mook 7b872cc41e [Perf] Skip bool-mask round-trip in block-sparse VSA attention (#1243) 2026-04-26 15:14:37 -07:00
Satyam Srivastava 3ff4a8d2d2 [ci] Refactor pr_test.sh post run hooks 2026-04-26 12:26:24 -07:00
Satyam Srivastava 9343d4cdf4 [bug] Fix modal remote functions crash container on sys.exit(0) during PR tests 2026-04-26 11:18:21 -07:00
Satyam Srivastava 66fb3d1e79 [ci] Fix buildkite agent annotate and upload 2026-04-26 02:36:45 -07:00
alexzms 37418946c8 [docs]: clarify real_score_guidance_scale CFG parameterization (#1256) 2026-04-26 16:38:00 +08:00
William Lin 95fd29e0cb [feat] Streaming WebSocket server skeleton (single generator + fMP4) (#1251) 2026-04-26 00:33:49 -07:00
Satyam Srivastava aca850cef2 [ci] Fix plotly bug 2026-04-25 20:17:49 -07:00
Satyam Srivastava 1c79779956 Debug changes 2026-04-25 18:15:09 -07:00
Satyam Srivastava eee03527ed [ci] Add plot render for performance tracking 2026-04-25 17:39:41 -07:00
Satyam Srivastava 1eb8541094 Test with different run configs 2026-04-25 15:53:36 -07:00
Junda Suandmergify[bot] e17cd2633c [bugfix]: normalize uint8 pil_image in I2V VAE encoding (#1249)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-24 09:16:01 +00:00
William Lin e0dc5f2b0c [feat] Add typed LTX-2 continuation state and streaming session store (#1250) 2026-04-24 01:28:07 -07:00
Satyam Srivastava 69c214d13a Test: HF performance sync logic 2026-04-23 16:06:04 -07:00
Satyam Srivastava 0341481aa7 [ci] Upload Perf. Regression Results to HF_Repo 2026-04-23 15:51:26 -07:00
William Lin 70ee5d230c [feat] [6/n] Improve API: LTX-2 public preset + asset wiring + gpu_pool translation (#1239) 2026-04-23 11:36:45 -07:00
Satyam Srivastava d1c3fdd187 [ci] Perf CI Run Schedule
Make Full Performance CI run only on main branch in X days and not on every PR.
This schedule will be created in the Buildkite.
2026-04-22 21:12:54 -07:00
Satyam Srivastava 980e8d933e Add CI Performance Regression Tracking Changes 2026-04-22 18:11:13 -07:00
William Lin 24ced500f5 [test] add LTX-2 distilled T2V SSIM regression test (#1240) 2026-04-21 12:03:38 -07:00
171 changed files with 13838 additions and 860 deletions
+96
View File
@@ -0,0 +1,96 @@
#!/usr/bin/env bash
# Sync .agents/skills/ into .claude/skills/ via per-skill symlinks.
#
# Why: Claude Code only scans .claude/skills/ and ~/.claude/skills/ for
# user-invocable skills (no skillsPath config exists — see
# https://code.claude.com/docs/en/skills.md). This repo's skills live
# in .agents/skills/ so they travel with the repo and stay under git.
# Run this once after cloning (or after adding/removing a skill) to
# expose them to Claude Code without maintaining a parallel tree.
#
# Usage:
# .agents/scripts/sync-skills.sh
#
# Idempotent and safe to re-run. Prunes stale symlinks whose source
# has been removed from .agents/skills/. Leaves hand-written
# .claude/skills/<name>/ directories untouched (only symlinks are
# managed).
set -euo pipefail
REPO_ROOT="$(git -C "$(dirname "$0")" rev-parse --show-toplevel)"
SRC_DIR="$REPO_ROOT/.agents/skills"
DST_DIR="$REPO_ROOT/.claude/skills"
if [[ ! -d "$SRC_DIR" ]]; then
echo "Error: $SRC_DIR does not exist." >&2
exit 1
fi
mkdir -p "$DST_DIR"
linked=0
unchanged=0
skipped=0
pruned=0
link_skill() {
local name="$1"
local src="$SRC_DIR/$name"
local dst="$DST_DIR/$name"
# Relative target keeps symlinks portable across clones.
local rel="../../.agents/skills/$name"
if [[ -L "$dst" ]]; then
if [[ "$(readlink "$dst")" == "$rel" ]]; then
unchanged=$((unchanged + 1))
return
fi
rm "$dst"
elif [[ -e "$dst" ]]; then
echo "Skipped (not a symlink): .claude/skills/$name" >&2
skipped=$((skipped + 1))
return
fi
ln -s "$rel" "$dst"
echo "Linked: .claude/skills/$name -> $rel"
linked=$((linked + 1))
}
prune_stale() {
local link="$1"
local target
target="$(readlink "$link")"
case "$target" in
../../.agents/skills/*) ;;
*) return ;;
esac
local name="${target##*/}"
if [[ ! -d "$SRC_DIR/$name" ]]; then
rm "$link"
echo "Pruned stale: .claude/skills/$(basename "$link")"
pruned=$((pruned + 1))
fi
}
for src in "$SRC_DIR"/*/; do
[[ -d "$src" ]] || continue
name="$(basename "$src")"
# Only treat directories that actually contain a SKILL.md as skills.
[[ -f "$src/SKILL.md" ]] || continue
link_skill "$name"
done
shopt -s nullglob
for link in "$DST_DIR"/*; do
[[ -L "$link" ]] || continue
prune_stale "$link"
done
shopt -u nullglob
printf "\nSummary: %d linked, %d unchanged, %d pruned" "$linked" "$unchanged" "$pruned"
if [[ "$skipped" -gt 0 ]]; then
printf ", %d skipped (non-symlink collision)" "$skipped"
fi
printf "\n"
+1
View File
@@ -5,3 +5,4 @@
{"name": "evaluate-video-quality", "description": "Evaluate generated video quality using available metrics (SSIM, loss trajectory, caption consistency)", "path": "evaluate-video-quality/SKILL.md", "status": "draft", "trust": "low"}
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
@@ -0,0 +1,250 @@
---
name: seed-ssim-references
description: Seed HF reference videos for a single newly-added SSIM test. Runs the test on Modal L40S, downloads the generated mp4s via `modal volume get`, pauses for the user to eyeball quality, then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
---
# Seed SSIM Reference Videos
## Purpose
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
reference videos exist on the HF dataset (`FastVideo/ssim-reference-videos`).
This skill:
1. Runs the test on Modal's L40S pool to generate the videos.
2. Downloads them to the local repo via `modal volume get`.
3. Pauses so the user can eyeball the mp4s and confirm quality.
4. Uploads only the new test's files to HF, with a guard that refuses to
overwrite anything already present.
The skill is run **manually**, once per new test. Before invoking it, the user
has already sanity-tested the new test locally — it launches `VideoGenerator`
and writes an mp4 without crashing. The skill does not re-test locally; it
goes straight to Modal L40S (which is what CI uses).
## When to use
- A new `test_*_similarity.py` file has been added in `fastvideo/tests/ssim/`
and the HF dataset has no `reference_videos/default/L40S_reference_videos/<model_id>/`
subtree for it yet.
## When not to use
- Regular CI runs — once refs exist, `pytest fastvideo/tests/ssim/` downloads
them automatically.
- Re-seeding an existing test. That requires `--force` on the upload step, and
is out of scope here; treat as a separate, deliberate operation.
## Inputs
The skill has **one required input**: the path to the new SSIM test file.
Prompt the user for it if they didn't supply it.
| Parameter | Required | Description |
|-----------|----------|-------------|
| `test_file` | Yes | e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`. The skill's first action is to ask for this if missing. |
Everything else is fixed:
- Modal runner GPU: **L40S** (hardcoded in `fastvideo/tests/modal/ssim_test.py`).
- Device folder: `L40S_reference_videos`.
- Quality tier: `default` (the tier CI runs). The `full_quality` tier is not
seeded by this skill.
- HF repo: `FastVideo/ssim-reference-videos` (dataset).
- Multi-model test files: all model ids in `*_MODEL_TO_PARAMS` are seeded
together; the Modal run produces one mp4 per (model, prompt, backend) and
the upload scopes by `--model-id`, looping if there is more than one.
## Prerequisites
The user has confirmed:
- `modal` CLI authenticated.
- `HF_API_KEY` (or `HUGGINGFACE_HUB_TOKEN` / `HF_TOKEN`) exported with write
access to `FastVideo/ssim-reference-videos`.
- The test file runs locally end-to-end (generates an mp4; SSIM assertion
failure due to missing reference is expected and fine).
Fail fast if the token env var is missing.
## Steps
### 1. Ask for the test file
If the user didn't name one, ask: *"Which SSIM test file do you want to seed
references for? (e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`)"*.
Validate:
- Path exists and matches `fastvideo/tests/ssim/test_*_similarity.py`.
- File defines a `*_MODEL_TO_PARAMS` dict — grep it to extract the set of
model ids. Those ids drive step 5.
If either check fails, stop and tell the user what's wrong.
### 2. Run the test on Modal L40S
Pick a subdir name so repeated runs don't collide:
```bash
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
```
Then launch the Modal run:
```bash
modal run fastvideo/tests/modal/ssim_test.py \
--git-repo="$(git config --get remote.origin.url)" \
--git-commit="$(git rev-parse HEAD)" \
--hf-api-key="$HF_API_KEY" \
--test-files="<test_file>" \
--sync-generated-to-volume \
--generated-volume-subdir="$SUBDIR" \
--skip-reference-download \
--no-fail-fast
```
Flag rationale:
- `--skip-reference-download`: no refs exist yet, so conftest must not try to
pull them.
- `--no-fail-fast`: lets the test finish generation before `_assert_similarity`
raises `FileNotFoundError: Reference video folder does not exist`. The
expected failure is what we want — the mp4 has already been written.
- `--sync-generated-to-volume` + `--generated-volume-subdir`: copies the
generated mp4s to the `hf-model-weights` Modal volume under
`ssim_generated_videos/default/<SUBDIR>/generated_videos/` so we can pull
them locally.
The Modal run will end with a nonzero exit (expected) and print a
`modal volume get hf-model-weights ssim_generated_videos/default/<SUBDIR>/generated_videos ./generated_videos_modal/default`
command. Capture that `<SUBDIR>` — you need it for step 3.
### 3. Download generated videos locally
```bash
modal volume get --force hf-model-weights \
ssim_generated_videos/default/"$SUBDIR"/generated_videos \
./generated_videos_modal/default
```
`--force` is required when the parent `./generated_videos_modal/default`
already exists; without it, `modal volume get` errors with `[Errno 21] Is a
directory`. Safe to pass on the first run too.
After this, the mp4s live at
`./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
The extra `generated_videos/` level comes from the volume layout in
`_sync_generated_videos_to_volume` (`ssim_test.py`) — the command copies
`<repo>/fastvideo/tests/ssim/generated_videos/<tier>` to
`ssim_generated_videos/<tier>/<SUBDIR>/generated_videos/`, and `modal volume
get` preserves that trailing `generated_videos/` segment.
### 4. PAUSE — user reviews quality
Print the list of downloaded mp4s and their paths, then stop. Tell the user:
> "Generated videos downloaded to `./generated_videos_modal/default/generated_videos/L40S_reference_videos/`. Please open them and confirm the quality looks correct. Reply **`upload`** to continue, or anything else to abort."
Do not proceed until the user explicitly says `upload`. If they abort, leave
everything on disk so they can inspect further — no cleanup.
### 5. Copy into the local reference layout
Scoped copy — only the new test's mp4s. Loop over each `<model_id>` extracted
in step 1:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
--quality-tier default \
--device-folder L40S_reference_videos \
--generated-dir ./generated_videos_modal/default/generated_videos/L40S_reference_videos
```
(The `--generated-dir` points at the device-folder root inside the
downloaded tree; `copy-local` walks all `<model>/<backend>/*.mp4`
underneath it. Since the Modal run was scoped to a single test file via
`--test-files`, only that test's model(s) are present — so the copy is
implicitly per-test.)
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
### 6. Upload to HF — scoped per model_id, with overwrite guard
For each `<model_id>`:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py upload \
--quality-tier default \
--device-folder L40S_reference_videos \
--model-id "<model_id>"
```
The upload command:
- Uploads **only** `reference_videos/default/L40S_reference_videos/<model_id>/`.
- **Refuses** if any file already exists at that path on HF (this is the
guard — seeding a new test should never clobber existing refs). To override,
the user must re-run with `--force`. If the guard fires, stop and report
exactly which files exist; do not silently `--force`.
Reads the HF token from `HF_API_KEY` / `HUGGINGFACE_HUB_TOKEN` / `HF_TOKEN`.
### 7. Report success
List what was uploaded (paths in repo) and remind the user to push any
related code changes. Do **not** auto-verify by re-running Modal — the user
can run `pytest fastvideo/tests/ssim/<test_file>` later to confirm end-to-end;
it will auto-download the refs they just uploaded.
## Failure modes and how to handle them
- **`HF_API_KEY` unset.** Stop before step 2. The Modal run needs it (passed
via `--hf-api-key`), and step 6 needs it for upload.
- **Modal run fails before generation.** No mp4s on the volume — nothing to
download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
and retry from step 2.
- **`./generated_videos_modal/default/L40S_reference_videos/` missing after
`modal volume get`.** The run didn't produce videos (most likely the test
crashed before writing, or `REQUIRED_GPUS` exceeded the partition capacity
— see Modal logs).
- **Upload guard fires (files already exist).** The test name / model id
collides with something already on HF. Verify the user actually wants to
replace existing refs; if so, re-run the upload with `--force`. If not,
rename the model id in `*_MODEL_TO_PARAMS` and re-seed.
- **Quality looks wrong in step 4.** Abort. The mp4s stay on disk for
inspection. The fix is usually in the test's params (resolution, steps,
seed) — edit the test, then re-run the skill.
## Design notes (for future skill maintainers)
- The skill deliberately runs on Modal, **not** locally, because the CI
runner is L40S. Seeding from a different GPU SKU produces refs that CI's
L40S runs can't match (SSIM drifts across SKUs).
- The skill is default-tier only. `full_quality` refs are seeded by a
separate, deliberate operation — they double runtime and aren't what CI
gates on.
- The overwrite guard in `reference_videos_cli.py upload` is default-on
specifically because this skill exists. Re-seeding is a distinct operation
that requires explicit `--force`.
## References
- `fastvideo/tests/modal/ssim_test.py` — Modal orchestrator; see
`--sync-generated-to-volume`, `--generated-volume-subdir`,
`--skip-reference-download`, `--no-fail-fast`.
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
(with `--model-id`, `--force`), `download`, `ensure` subcommands.
- `fastvideo/tests/ssim/README.md` — reference layout, HF repo conventions.
- `fastvideo/tests/ssim/inference_similarity_utils.py` —
`run_text_to_video_similarity_test` + `_build_init_kwargs`: what each test
config passes to `VideoGenerator.from_pretrained`.
## Changelog
| Date | Change |
|------|--------|
| 2026-04-17 | Initial version (Modal sync-to-volume flow). |
| 2026-04-21 | Rewrite: single-test scope, explicit user-review pause, per-`model_id` upload, HF overwrite guard. Dropped `scripts/seed_ssim.sh`. |
| 2026-04-21 | Post-first-run fixes: `modal volume get` needs `--force` when parent exists; download tree has an extra `generated_videos/` level so `--generated-dir` must reflect it. |
@@ -29,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": 1,
"num_measurement_runs": 3,
"num_warmup_runs": 2,
"num_measurement_runs": 5,
"required_gpus": 2
},
"thresholds": {
+73 -2
View File
@@ -63,7 +63,72 @@ 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 IMAGE_VERSION=$IMAGE_VERSION"
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
}
case "$TEST_TYPE" in
"encoder")
@@ -124,8 +189,9 @@ 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..."
log "Running performance tests on Modal..."
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..."
@@ -147,5 +213,10 @@ 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,6 +18,7 @@ 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
+6 -3
View File
@@ -296,8 +296,10 @@ 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/` (e.g., `tests/local_tests/upsamplers/`)
and the commands in `tests/local_tests/README.md`.
- 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`.
- 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`).
@@ -348,7 +350,8 @@ Purpose:
Action:
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
- Add a pipeline parity test under `tests/local_tests/<family>/`
(e.g., `tests/local_tests/<family>/test_<family>_pipeline_parity.py`).
- See the [Testing Guide](testing.md) for test conventions.
### 7) Add user‑facing examples
@@ -29,7 +29,7 @@ surfaces:
vae_cpu_offload: generator.engine.offload.vae
pin_cpu_memory: generator.engine.offload.pin_cpu_memory
enable_torch_compile: generator.engine.compile.enabled
torch_compile_kwargs: generator.engine.compile.kwargs
torch_compile_kwargs: generator.engine.compile.backend,fullgraph,mode,dynamic,extras
disable_autocast: generator.engine.disable_autocast
enable_stage_verification: generator.engine.enable_stage_verification
prompt_txt: request.inputs.prompt_path
@@ -40,8 +40,8 @@ surfaces:
init_weights_from_safetensors_2: generator.pipeline.components.transformer_2_weights
override_pipeline_cls_name: generator.pipeline.components.override_pipeline_cls_name
boundary_ratio: request.sampling.boundary_ratio
ltx2_vae_tiling: generator.pipeline.vae_tiling
preset_owned:
ltx2_vae_tiling: generator.pipeline.preset_overrides.ltx2.vae_tiling
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.preset_overrides.ltx2.vae.spatial_tile_size_in_pixels
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.preset_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.preset_overrides.ltx2.vae.temporal_tile_size_in_frames
@@ -354,6 +354,8 @@ surfaces:
return_frames: request.output.return_frames
return_trajectory_latents: request.runtime.return_trajectory_latents
return_trajectory_decoded: request.runtime.return_trajectory_decoded
continuation_state: request.state
return_continuation_state: request.output.return_state
preset_owned:
t_thresh: request.stage_overrides.refine.t_thresh
spatial_refine_only: request.stage_overrides.refine.spatial_refine_only
+177
View File
@@ -0,0 +1,177 @@
# Streaming WebSocket Server Contract
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
document is the authoritative spec for the message catalogue and the
session state machine. Any change to either must update this document
in the same PR that touches `protocol.py` or `session.py`.
## Endpoint
| Path | Protocol | Purpose |
|---|---|---|
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
The server is launched by `fastvideo serve --config <serve.yaml>` when
the config carries a `streaming:` block. Without that block the same CLI
launches the OpenAI stateless HTTP server instead.
## Connection lifecycle
Every WebSocket connection holds exactly one `Session`. Sessions move
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
```
┌──────────────┐
│ INITIALIZING │ ← WebSocket accepted, before init frame
└──────┬───────┘
│ session_init_v2 received
┌──────────────┼──────────────┐
▼ ▼ ▼
QUEUED GPU_BINDING REJECTED
│ │ ↑
│ slot ready │ │ max-sessions hit
▼ ▼ │ or invalid init
┌────────┐ │
│ ACTIVE │ ────────┘
└────┬───┘
segment loop │
│
┌───────────┼───────────┐
▼ ▼ ▼
COMPLETE ERROR TIMEOUT
(clean leave) (any failure) (idle / segment_cap reached)
```
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
no transitions out. The transition matrix is enforced in
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
`SessionManager` enforces the per-process budgets pulled from
`StreamingConfig`:
- `session_timeout_seconds` — idle reaper drops sessions that haven't
advanced; non-terminal sessions transition to `TIMEOUT`.
- `generation_segment_cap` — a session that hits the cap transitions to
`COMPLETE` after the last segment ships.
## Message catalogue
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
`protocol.py` are the source of truth; this table is the human-readable
view.
### Client → server
| `type` | Required fields | Purpose |
|---|---|---|
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
The opening frame must be `session_init_v2`. Any other first frame is
rejected with an `error` (code `invalid_message`) and the WebSocket is
closed.
### Server → client
| `type` | Carries | When emitted |
|---|---|---|
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
## Continuation state
The session optionally accepts a `continuation_state` dict inside the
opening `session_init_v2` frame. When present, the server hydrates it
into a `ContinuationState(kind, payload)` envelope and feeds it as the
`request.state` on the first segment's `GenerationRequest` — letting a
client resume after a disconnect, migrate sessions across processes,
or replay a prior session.
After every segment, if the runtime returns a fresh state, the server
persists it to the `SessionStore` so a `snapshot_state` request can
export it. The store and serialization contracts live with the model
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
## Example flow
```
client server
────── ──────
WS /v1/stream ─────── connect ─────────────────────────►
◄────── (accept)
{"type": "session_init_v2",
"preset": "ltx2_two_stage",
"curated_prompts": ["a fox in snow", "the fox jumps"],
"initial_image": {...},
"stream_mode": "av_fmp4"} ─────────────────────────────►
(validate, queue, bind)
◄──── {"type": "queue_status",
"position": 0, "queue_depth": 0}
◄──── {"type": "gpu_assigned",
"gpu_id": 0, "model_id": "..."}
◄──── {"type": "ltx2_stream_start", ...}
{"type": "segment_prompt_source",
"prompt": "a fox in snow",
"source": "curated"} ───────────────────────────────────►
(run pipeline)
◄──── {"type": "ltx2_segment_start",
"segment_idx": 1, ...}
◄──── {"type": "step_complete",
"segment_idx": 1, "timings": {...}}
◄──── {"type": "media_init",
"segment_idx": 1,
"mime": "video/mp4", ...}
◄──── <binary fMP4 init segment>
◄──── <binary fMP4 fragment>
◄──── <binary fMP4 fragment>
◄──── {"type": "media_segment_complete",
"segment_idx": 1, "chunks": 12}
◄──── {"type": "ltx2_segment_complete",
"segment_idx": 1, ...}
{"type": "segment_prompt_source",
"prompt": "the fox jumps"} ─────────────────────────────►
(segment 2 …)
{"type": "snapshot_state"} ──────────────────────────────►
◄──── {"type": "continuation_state_snapshot",
"kind": "ltx2.v1",
"payload": {"schema_version": 1, ...}}
(close) ──────────────────────────────────────────────────►
(session → COMPLETE)
```
## Backward / forward compatibility
- Adding a new client message: append a Pydantic model to `protocol.py`
with a unique `type`; add the discriminator entry to `ClientMessage`;
add a row to the table above. Old clients that don't send the new
message remain compatible.
- Adding a new server message: emit only when a new feature flag is
enabled (or always emit, since clients ignore unknown types).
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
→ `session_init_v3`) and accept both for one release cycle. Never
silently change field semantics under the same `type`.
+22
View File
@@ -86,3 +86,25 @@ sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
- Learning rate: 2e-5
- Training steps: 3000 (~12 hours)
- HSDP shard dim: 1
## 🧭 Note on `real_score_guidance_scale`
The teacher CFG used inside the DMD loss follows the DMD2 reference
implementation and uses the parameterization
```
x = x_cond + w * (x_cond - x_uncond)
```
rather than the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)`. The
two are mathematically equivalent up to a constant offset:
| `real_score_guidance_scale` (`w`) | Equivalent standard CFG (`w + 1`) | Output |
|-----------------------------------|-----------------------------------|-----------------------|
| `-1` | `0` | unconditional |
| `0` | `1` | conditional |
| `3.5` (default) | `4.5` | strong guidance |
So `real_score_guidance_scale` should be read as the **extra** guidance
strength added on top of the conditional prediction. When porting values
from a paper that uses the Ho & Salimans form, subtract 1.
@@ -0,0 +1,77 @@
# 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()
@@ -0,0 +1,77 @@
# 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()
@@ -0,0 +1,84 @@
# 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()
@@ -0,0 +1,53 @@
# 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()
@@ -0,0 +1,73 @@
# 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
@@ -0,0 +1,79 @@
# 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
+10 -1
View File
@@ -61,7 +61,16 @@ has_cmake_arg() {
}
detect_with_torch() {
uv run --active --no-project python -c "import torch
# Prefer the active venv's python directly over `uv run --active --no-project`,
# which on some uv versions provisions its own interpreter and misses packages
# installed into VIRTUAL_ENV.
local py
if [[ -n "${VIRTUAL_ENV:-}" && -x "${VIRTUAL_ENV}/bin/python" ]]; then
py="${VIRTUAL_ENV}/bin/python"
else
py="$(command -v python3 || command -v python)"
fi
"${py}" -c "import torch
if not torch.cuda.is_available():
raise RuntimeError('torch.cuda.is_available() is false')
mj, mn = torch.cuda.get_device_capability(0)
@@ -5,6 +5,11 @@ from fastvideo_kernel.ops import (
video_sparse_attn,
)
from fastvideo_kernel.block_sparse_attn import (
block_sparse_attn,
block_sparse_attn_from_indices,
)
from fastvideo_kernel.vmoba import (
moba_attn_varlen,
process_moba_input,
@@ -22,6 +27,8 @@ from fastvideo_kernel.turbodiffusion_ops import (
__all__ = [
"sliding_tile_attention",
"video_sparse_attn",
"block_sparse_attn",
"block_sparse_attn_from_indices",
"moba_attn_varlen",
"process_moba_input",
"process_moba_output",
@@ -1,3 +1,5 @@
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
from __future__ import annotations
import os
@@ -6,6 +8,11 @@ from typing import Tuple
import torch
# ---------------------------------------------------------------------------
# Backend selection helpers
# ---------------------------------------------------------------------------
def _get_sm90_ops():
try:
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
@@ -25,38 +32,66 @@ def _is_sm90() -> bool:
def _force_triton() -> bool:
# Force Triton even on SM90 and even if the compiled extension is available.
# Useful for CI / debugging / parity testing.
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Preferred map->index conversion used by the wrapper.
# ---------------------------------------------------------------------------
# Index helpers
# ---------------------------------------------------------------------------
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
"""
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Compact a bool block_map to (q2k_idx, q2k_num). Legacy path only."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
if block_map.dim() != 4:
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
raise ValueError(
f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), "
f"got shape={tuple(block_map.shape)}"
)
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
if not block_map.is_cuda:
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
raise RuntimeError(
"block_map must be a CUDA tensor (Triton map_to_index required)."
)
try:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index
except Exception as e: # pragma: no cover - environment issue
raise ImportError(
"Triton map_to_index is required but not available. "
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
"Ensure Triton is installed and "
"fastvideo_kernel.triton_kernels.index is importable."
) from e
return triton_map_to_index(block_map)
def _invert_indices_for_backward(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
from fastvideo_kernel.triton_kernels.index import invert_indices
return invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
"""Return `t` as a contiguous int32 tensor, raising a clear error on CPU input."""
if not t.is_cuda:
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
if t.dtype != torch.int32:
t = t.to(torch.int32)
if not t.is_contiguous():
t = t.contiguous()
return t
# ---------------------------------------------------------------------------
# Triton backend custom ops (index-native)
# ---------------------------------------------------------------------------
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_triton",
mutates_args=(),
@@ -66,34 +101,40 @@ def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
triton_block_sparse_attn_forward,
)
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
o, M = triton_block_sparse_attn_forward(
q.contiguous(),
k.contiguous(),
v.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
M = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
device=q.device,
dtype=torch.float32,
)
return o, M
@@ -109,20 +150,32 @@ def block_sparse_attn_backward_triton(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
triton_block_sparse_attn_backward,
)
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
# o/M are kernel outputs so are already contiguous.
dq, dk, dv = triton_block_sparse_attn_backward(
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
grad_output.contiguous(),
q.contiguous(),
k.contiguous(),
v.contiguous(),
o,
M,
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
variable_block_sizes,
)
return dq, dk, dv
@@ -135,7 +188,8 @@ def _block_sparse_attn_backward_triton_fake(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q)
@@ -144,19 +198,28 @@ def _block_sparse_attn_backward_triton_fake(
return dq, dk, dv
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def _setup_context_triton(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
block_sparse_attn_triton.register_autograd(
_backward_triton, setup_context=_setup_context_triton
)
# ---------------------------------------------------------------------------
# SM90 backend custom ops (index-native)
# ---------------------------------------------------------------------------
@torch.library.custom_op(
@@ -168,21 +231,21 @@ def block_sparse_attn_sm90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
block_sparse_fwd, _ = _get_sm90_ops()
if block_sparse_fwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o_padded, lse_padded
@@ -192,11 +255,16 @@ def _block_sparse_attn_sm90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q_padded)
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
lse = torch.empty(
(q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1),
device=q_padded.device,
dtype=torch.float32,
)
return o, lse
@@ -212,30 +280,34 @@ def block_sparse_attn_backward_sm90(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
_, block_sparse_bwd = _get_sm90_ops()
if block_sparse_bwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
dq, dk, dv = block_sparse_bwd(
q_padded,
k_padded,
v_padded,
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
o_padded,
lse_padded,
grad_output_padded,
grad_output_padded.contiguous(),
k2q_idx,
k2q_num,
variable_block_sizes.int(),
variable_block_sizes,
)
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
# C++ kernel returns fp32 grads; cast back to the input dtype.
out_dtype = grad_output_padded.dtype
return dq.to(out_dtype), dk.to(out_dtype), dv.to(out_dtype)
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
@@ -246,7 +318,8 @@ def _block_sparse_attn_backward_sm90_fake(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q_padded)
@@ -255,21 +328,57 @@ def _block_sparse_attn_backward_sm90_fake(
return dq, dk, dv
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
)
return dq, dk, dv, None, None
def _setup_context_sm90(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, lse = output
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
block_sparse_attn_sm90.register_autograd(
_backward_sm90, setup_context=_setup_context_sm90
)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def block_sparse_attn_from_indices(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Block-sparse attention with autograd, taking compact per-row KV indices."""
# Normalize index tensors once at the public boundary so the custom ops
# and their fakes can assume int32/contiguous. No-op on well-formed input.
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
variable_block_sizes = _as_int32_contig(variable_block_sizes, "variable_block_sizes")
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
use_sm90 = (
(not _force_triton())
and _is_sm90()
and block_sparse_fwd is not None
and block_sparse_bwd is not None
)
if use_sm90:
return block_sparse_attn_sm90(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
def block_sparse_attn(
@@ -279,16 +388,8 @@ def block_sparse_attn(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Unified block-sparse attention op with autograd support.
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
"""
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
"""Bool-mask compat wrapper; prefer block_sparse_attn_from_indices."""
q2k_idx, q2k_num = _map_to_index(block_map)
return block_sparse_attn_from_indices(
q, k, v, q2k_idx, q2k_num, variable_block_sizes
)
@@ -1,6 +1,6 @@
import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
# Try to load the C++ extension
@@ -125,13 +125,18 @@ def video_sparse_attn(
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, q_seq_len, dim)
# Sparse branch
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
topk_idx = torch.topk(scores, topk, dim=-1).indices
mask = torch.zeros_like(scores,
dtype=torch.bool).scatter_(-1, topk_idx, True)
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
q2k_idx = topk_idx.to(torch.int32).contiguous()
q2k_num = torch.full(
(batch, heads, q_num_blocks),
topk,
dtype=torch.int32,
device=q.device,
)
out_s = block_sparse_attn_from_indices(
q, k, v, q2k_idx, q2k_num, variable_block_sizes
)[0]
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
@@ -1,9 +1,10 @@
## pytorch sdpa version of block sparse ##
from typing import Tuple
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
@@ -153,3 +154,114 @@ def map_to_index(block_map: torch.Tensor):
)
return index, index_num
@triton.jit
def _invert_indices_kernel(
q2k_idx_ptr,
q2k_num_ptr,
k2q_idx_ptr,
k2q_num_ptr,
q2k_idx_b, q2k_idx_h, q2k_idx_q, q2k_idx_k,
q2k_num_b, q2k_num_h, q2k_num_q,
k2q_idx_b, k2q_idx_h, k2q_idx_k, k2q_idx_q,
k2q_num_b, k2q_num_h, k2q_num_k,
MAX_KV_PER_Q: tl.constexpr,
):
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
pid_q = tl.program_id(2)
n = tl.load(
q2k_num_ptr
+ pid_b * q2k_num_b
+ pid_h * q2k_num_h
+ pid_q * q2k_num_q
)
q2k_row = (
q2k_idx_ptr
+ pid_b * q2k_idx_b
+ pid_h * q2k_idx_h
+ pid_q * q2k_idx_q
)
for i in tl.range(0, MAX_KV_PER_Q):
if i < n:
kv = tl.load(q2k_row + i * q2k_idx_k)
count_ptr = (
k2q_num_ptr
+ pid_b * k2q_num_b
+ pid_h * k2q_num_h
+ kv * k2q_num_k
)
pos = tl.atomic_add(count_ptr, 1)
tl.store(
k2q_idx_ptr
+ pid_b * k2q_idx_b
+ pid_h * k2q_idx_h
+ kv * k2q_idx_k
+ pos * k2q_idx_q,
pid_q,
)
def invert_indices(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
if q2k_idx.dim() != 4:
raise ValueError(
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
)
if q2k_num.dim() != 3:
raise ValueError(
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
)
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
raise RuntimeError("invert_indices requires CUDA tensors.")
B, H, Nq, Mk = q2k_idx.shape
if q2k_num.shape != (B, H, Nq):
raise ValueError(
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
f"[B, H, Nq] = {(B, H, Nq)}"
)
q2k_idx = q2k_idx.contiguous()
q2k_num = q2k_num.contiguous()
if q2k_idx.dtype != torch.int32:
q2k_idx = q2k_idx.to(torch.int32)
if q2k_num.dtype != torch.int32:
q2k_num = q2k_num.to(torch.int32)
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
# `Nq` is a tight upper bound on the compacted K->Q slots.
k2q_idx = torch.empty(
(B, H, num_kv_blocks, Nq),
dtype=torch.int32,
device=q2k_idx.device,
)
k2q_num = torch.zeros(
(B, H, num_kv_blocks),
dtype=torch.int32,
device=q2k_idx.device,
)
grid = (B, H, Nq)
_invert_indices_kernel[grid](
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
q2k_idx.stride(0), q2k_idx.stride(1), q2k_idx.stride(2), q2k_idx.stride(3),
q2k_num.stride(0), q2k_num.stride(1), q2k_num.stride(2),
k2q_idx.stride(0), k2q_idx.stride(1), k2q_idx.stride(2), k2q_idx.stride(3),
k2q_num.stride(0), k2q_num.stride(1), k2q_num.stride(2),
MAX_KV_PER_Q=Mk,
)
return k2q_idx, k2q_num
+118 -9
View File
@@ -16,6 +16,8 @@ from fastvideo.api.request_metadata import (
reset_tracking_roots,
)
from fastvideo.api.schema import (
CompileConfig,
ContinuationState,
GenerationRequest,
GeneratorConfig,
InputConfig,
@@ -25,6 +27,10 @@ from fastvideo.api.schema import (
)
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_preset_override_fields,
refine_stage_override_fields,
)
from fastvideo.utils import shallow_asdict
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
@@ -38,6 +44,10 @@ _LEGACY_REQUEST_ALIASES = {
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
"embedded_cfg_scale",
})
# torch.compile kwargs that map to first-class CompileConfig fields.
_COMPILE_TYPED_KEYS = ("backend", "fullgraph", "mode", "dynamic")
# LTX-2 refine flat kwargs (init + per-request) known to FastVideoArgs.
_LTX2_REFINE_FLAT_KEYS = (refine_preset_override_fields() | refine_stage_override_fields())
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
@@ -80,6 +90,8 @@ def legacy_from_pretrained_to_config(
components: dict[str, Any] = {}
quantization: dict[str, Any] = {}
experimental: dict[str, Any] = {}
preset_overrides: dict[str, Any] = {}
preset_refine: dict[str, Any] = {}
for key, value in kwargs.items():
if key == "revision":
@@ -106,8 +118,33 @@ def legacy_from_pretrained_to_config(
offload["pin_cpu_memory"] = value
elif key == "enable_torch_compile":
compile_config["enabled"] = value
elif key == "enable_torch_compile_text_encoder":
compile_config["text_encoder_enabled"] = value
elif key == "torch_compile_kwargs":
compile_config["kwargs"] = deepcopy(value)
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
for first_class in _COMPILE_TYPED_KEYS:
if first_class in remaining:
compile_config[first_class] = remaining.pop(first_class)
if remaining:
compile_config["extras"] = remaining
elif key == "ltx2_vae_tiling":
pipeline["vae_tiling"] = value
elif key == "config_model_path":
components["config_root"] = value
elif key == "ltx2_refine_enabled":
preset_refine["enabled"] = value
elif key == "ltx2_refine_upsampler_path":
# Empty string means "no upsampler"; keep typed None.
components["upsampler_weights"] = value or None
elif key == "ltx2_refine_lora_path":
# Empty string means "no refine LoRA"; keep typed None.
components["lora_path"] = value or None
elif key == "ltx2_refine_add_noise":
preset_refine["add_noise"] = value
elif key == "ltx2_refine_num_inference_steps":
preset_refine["num_inference_steps"] = value
elif key == "ltx2_refine_guidance_scale":
preset_refine["guidance_scale"] = value
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
engine[key] = value
elif key == "override_text_encoder_quant":
@@ -147,6 +184,10 @@ def legacy_from_pretrained_to_config(
if components:
pipeline["components"] = components
if preset_refine:
preset_overrides["refine"] = preset_refine
if preset_overrides:
pipeline["preset_overrides"] = preset_overrides
if experimental:
pipeline["experimental"] = experimental
if pipeline:
@@ -162,12 +203,8 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
unsupported.append("pipeline.preset")
if normalized.pipeline.preset_version is not None:
unsupported.append("pipeline.preset_version")
if normalized.pipeline.components.config_root is not None:
unsupported.append("pipeline.components.config_root")
if normalized.pipeline.components.vae_weights is not None:
unsupported.append("pipeline.components.vae_weights")
if normalized.pipeline.components.upsampler_weights is not None:
unsupported.append("pipeline.components.upsampler_weights")
if unsupported:
joined = ", ".join(unsupported)
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
@@ -191,13 +228,21 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
"vae_cpu_offload": engine.offload.vae,
"pin_cpu_memory": engine.offload.pin_cpu_memory,
"enable_torch_compile": engine.compile.enabled,
"torch_compile_kwargs": deepcopy(engine.compile.kwargs),
"torch_compile_kwargs": _compile_config_to_torch_kwargs(engine.compile),
"enable_stage_verification": engine.enable_stage_verification,
"use_fsdp_inference": engine.use_fsdp_inference,
"disable_autocast": engine.disable_autocast,
}
if normalized.pipeline.workload_type is not None:
kwargs["workload_type"] = normalized.pipeline.workload_type
if normalized.pipeline.vae_tiling is not None:
kwargs["ltx2_vae_tiling"] = normalized.pipeline.vae_tiling
if engine.compile.text_encoder_enabled is not None:
# ``FastVideoArgs.from_kwargs`` filters to declared fields, so
# this is a no-op on the current legacy path. Emit anyway so the
# realtime runtime (PR 7.6) — which reads from the kwargs dict
# before FastVideoArgs filtering — can pick it up once wired.
kwargs["enable_torch_compile_text_encoder"] = (engine.compile.text_encoder_enabled)
quantization = engine.quantization
if quantization is not None and quantization.text_encoder_quant is not None:
@@ -220,8 +265,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
kwargs["init_weights_from_safetensors"] = components.transformer_weights
if components.transformer_2_weights is not None:
kwargs["init_weights_from_safetensors_2"] = components.transformer_2_weights
if components.config_root is not None:
kwargs["config_model_path"] = components.config_root
if components.upsampler_weights is not None:
kwargs["ltx2_refine_upsampler_path"] = components.upsampler_weights
kwargs.update(deepcopy(normalized.pipeline.preset_overrides))
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
refine = preset_overrides.pop("refine", None)
if isinstance(refine, Mapping):
for key in _LTX2_REFINE_FLAT_KEYS:
if key in refine:
kwargs[f"ltx2_refine_{key}"] = refine[key]
kwargs.update(preset_overrides)
kwargs.update(deepcopy(normalized.pipeline.experimental))
return FastVideoArgs.from_kwargs(**kwargs)
@@ -271,10 +326,13 @@ def request_to_sampling_param(
) -> SamplingParam:
if request.plan is not None:
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
if request.state is not None:
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
sampling_param = SamplingParam.from_pretrained(model_path)
if request.state is not None:
_validate_continuation_state(request.state)
sampling_param.continuation_state = request.state
if request.output.return_state:
sampling_param.return_continuation_state = True
updates = explicit_request_updates(request)
for key, value in updates.items():
@@ -316,6 +374,25 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
return isinstance(raw.get("generator"), Mapping)
def _compile_config_to_torch_kwargs(compile_config: CompileConfig, ) -> dict[str, Any]:
"""Flatten typed ``CompileConfig`` back to a ``torch_compile_kwargs``
dict that the legacy ``FastVideoArgs`` path still expects.
Typed first-class fields (:attr:`backend`, :attr:`fullgraph`,
:attr:`mode`, :attr:`dynamic`) are only emitted when the user set
them explicitly (non-``None``). ``extras`` is merged on top for any
uncommon kwargs.
"""
out: dict[str, Any] = {}
for key in _COMPILE_TYPED_KEYS:
value = getattr(compile_config, key)
if value is not None:
out[key] = value
if compile_config.extras:
out.update(deepcopy(compile_config.extras))
return out
def _sampling_param_to_request_raw(sampling_param: SamplingParam | None, ) -> dict[str, Any]:
if sampling_param is None:
return {}
@@ -476,6 +553,37 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
_SCHEMA_DEFAULT_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
_KNOWN_CONTINUATION_KINDS: set[str] = set()
def register_continuation_kind(kind: str) -> None:
"""Register a :class:`ContinuationState.kind` as recognized.
PR 7 wires the envelope through; per-kind payload deserializers live
with each model family (e.g. ``fastvideo.pipelines.basic.ltx2.
continuation.LTX2ContinuationState``). The registry lets the
public-API compat layer validate the kind early, before the state
reaches the pipeline.
"""
if not isinstance(kind, str) or not kind:
raise ValueError("ContinuationState kind must be a non-empty string")
_KNOWN_CONTINUATION_KINDS.add(kind)
def _validate_continuation_state(state: ContinuationState) -> None:
if not isinstance(state.kind, str) or not state.kind:
raise ValueError("GenerationRequest.state.kind must be a non-empty string; got "
f"{state.kind!r}")
if not isinstance(state.payload, Mapping):
raise ValueError(f"GenerationRequest.state.payload must be a mapping; got "
f"{type(state.payload).__name__}")
if state.kind not in _KNOWN_CONTINUATION_KINDS:
known = sorted(_KNOWN_CONTINUATION_KINDS)
raise ValueError(f"Unknown ContinuationState kind {state.kind!r}; registered "
f"kinds: {known}. Import the model family that owns this kind "
"(e.g. `import fastvideo.pipelines.basic.ltx2.continuation`) "
"to register it, or drop the state field.")
def _fan_out_batched_input_value(
source_request: GenerationRequest,
@@ -509,6 +617,7 @@ __all__ = [
"load_generator_config_from_file",
"normalize_generation_request",
"normalize_generator_config",
"register_continuation_kind",
"request_to_pipeline_overrides",
"request_to_sampling_param",
]
+4
View File
@@ -15,6 +15,7 @@ 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
@@ -44,6 +45,7 @@ class GenerationResult:
"samples",
"frames",
"audio",
"audio_sample_rate",
"size",
"generation_time",
"logging_info",
@@ -62,6 +64,7 @@ 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"),
@@ -80,6 +83,7 @@ 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,
+48 -6
View File
@@ -1,11 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import copy
from dataclasses import dataclass, field, fields
from typing import Any
from typing import TYPE_CHECKING, Any
from fastvideo.logger import init_logger
from fastvideo.utils import StoreBoolean
if TYPE_CHECKING:
from fastvideo.api.schema import ContinuationState
logger = init_logger(__name__)
@@ -92,9 +97,13 @@ class SamplingParam:
movement_distance: float | None = None
camera_rotation: str | None = None
# LTX2 multi-modal CFG and STG
ltx2_cfg_scale_video: float = 3.0
ltx2_cfg_scale_audio: float = 7.0
# LTX-2 multi-modal CFG and STG.
# cfg_scale defaults are 1.0 (CFG off) so ``ForwardBatch.__post_init__``
# doesn't force ``do_classifier_free_guidance`` on non-LTX-2 models that
# never override these fields. LTX-2 presets that need text-CFG on set
# them in their ``defaults`` dict (e.g. ``ltx2_base``).
ltx2_cfg_scale_video: float = 1.0
ltx2_cfg_scale_audio: float = 1.0
ltx2_modality_scale_video: float = 3.0
ltx2_modality_scale_audio: float = 3.0
ltx2_rescale_scale: float = 0.7
@@ -103,6 +112,39 @@ 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
@@ -127,7 +169,7 @@ class SamplingParam:
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str) -> "SamplingParam":
def from_pretrained(cls, model_path: str) -> SamplingParam:
sampling_param = cls._from_preset(model_path)
if sampling_param is not None:
return sampling_param
@@ -143,7 +185,7 @@ class SamplingParam:
def _from_preset(
cls,
model_path: str,
) -> "SamplingParam | None":
) -> SamplingParam | None:
"""Build a SamplingParam from preset defaults.
Returns ``None`` when no preset is configured for
+20 -1
View File
@@ -33,8 +33,25 @@ class OffloadConfig:
@dataclass
class CompileConfig:
"""Typed ``torch.compile`` configuration.
``backend``/``fullgraph``/``mode``/``dynamic`` are the four most
common ``torch.compile`` knobs. ``extras`` holds any remaining
``torch.compile`` kwargs (e.g. ``options``, ``disable``).
"""
enabled: bool = False
kwargs: dict[str, Any] = field(default_factory=dict)
text_encoder_enabled: bool | None = None
"""Whether ``torch.compile`` is applied to the text encoder. ``None``
keeps the runtime default. The public ``FastVideoArgs`` adapter does
not yet consume this flag; reserved so the realtime runtime upstream
(PR 7.6) has a typed home for its ``enable_torch_compile_text_encoder``
kwarg without routing through ``pipeline.experimental``."""
backend: str | None = None
fullgraph: bool | None = None
mode: str | None = None
dynamic: bool | None = None
extras: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -76,6 +93,8 @@ class PipelineSelection:
preset: str | None = None
preset_version: int | None = None
components: ComponentConfig = field(default_factory=ComponentConfig)
vae_tiling: bool | None = None
"""Tile-based VAE decode. ``None`` keeps the model's default."""
preset_overrides: dict[str, Any] = field(default_factory=dict)
experimental: dict[str, Any] = field(default_factory=dict)
-2
View File
@@ -48,8 +48,6 @@ 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__()
+3 -1
View File
@@ -5,11 +5,13 @@ 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"
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"StableAudioConfig"
]
+2 -1
View File
@@ -50,7 +50,8 @@ class CosmosArchConfig(DiTArchConfig):
})
# Cosmos-specific config parameters based on transformer_cosmos.py
in_channels: int = 16
# in_channels includes the condition_mask channel (16 latent + 1 cond = 17)
in_channels: int = 17
out_channels: int = 16
num_attention_heads: int = 16
attention_head_dim: int = 128
@@ -0,0 +1,76 @@
# 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,9 +7,12 @@ 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"
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig"
]
@@ -0,0 +1,80 @@
# 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,6 +41,14 @@ 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,6 +5,7 @@ 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__ = [
@@ -16,4 +17,6 @@ __all__ = [
"Gen3CVAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
"OobleckVAEArchConfig",
"OobleckVAEConfig",
]
+68
View File
@@ -0,0 +1,68 @@
# 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"
+1 -1
View File
@@ -5,7 +5,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
from fastvideo.registry import get_pipeline_config_cls_from_name
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig, WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
+27 -2
View File
@@ -11,10 +11,34 @@ import torch
from fastvideo.configs.models import DiTConfig, VAEConfig
from fastvideo.configs.models.dits.base import DiTArchConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class LongCatT5ArchConfig(T5ArchConfig):
"""T5 arch that pads tokenizer output to ``max_length``.
LongCat's denoising stage concatenates positive and negative
attention masks along the batch dimension for CFG, which requires
uniform seq length. The shared :class:`T5ArchConfig` dropped the
``"padding": "max_length"`` tokenizer kwarg so other DiTs could run
with variable-length masks; LongCat still needs the uniform
contract.
"""
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LongCatT5Config(T5Config):
arch_config: TextEncoderArchConfig = field(default_factory=LongCatT5ArchConfig)
@dataclass
class LongCatDiTArchConfig(DiTArchConfig):
"""Extended DiTArchConfig with LongCat-specific fields."""
@@ -103,8 +127,9 @@ class LongCatT2V480PConfig(PipelineConfig):
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
# UMT5 uses T5-like config; postprocess pads to 512. LongCatT5Config
# restores ``padding="max_length"`` for the CFG concat contract.
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (LongCatT5Config(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (longcat_preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (umt5_postprocess_text, ))
@@ -0,0 +1,67 @@
# 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
+29 -2
View File
@@ -1,4 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.entrypoints.streaming.server import run_server
from fastvideo.entrypoints.streaming.server import build_app, run_server
from fastvideo.entrypoints.streaming.session import (
Session,
SessionManager,
SessionState,
)
from fastvideo.entrypoints.streaming.session_store import (
BlobStore,
InMemoryBlobStore,
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
)
__all__ = ["run_server"]
__all__ = [
"BlobStore",
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
"InMemoryBlobStore",
"InMemorySessionStore",
"Session",
"SessionManager",
"SessionState",
"SessionStore",
"build_app",
"run_server",
]
+252
View File
@@ -0,0 +1,252 @@
# SPDX-License-Identifier: Apache-2.0
"""JSON WebSocket protocol schemas for the streaming server.
Every control message shares the envelope ``{"type": <str>, ...}``.
Pydantic models live here so the server can parse / validate incoming
frames and emit well-typed outgoing frames without hand-rolled dicts.
The message catalogue matches the contract in
``docs/design/server_contracts/streaming.md``; additions must land in
both places in the same PR.
"""
from __future__ import annotations
from typing import Annotated, Any, Literal, Union
from pydantic import BaseModel, ConfigDict, Field
# ---------------------------------------------------------------------------
# Client → server
# ---------------------------------------------------------------------------
class SessionInitV2(BaseModel):
"""Opening frame the client sends after the WebSocket handshake."""
model_config = ConfigDict(extra="allow")
type: Literal["session_init_v2"]
client_id: str | None = None
preset: str | None = None
preset_label: str | None = None
curated_prompts: list[str] = Field(default_factory=list)
initial_image: dict[str, Any] | None = None
enhancement_enabled: bool = False
auto_extension_enabled: bool = False
loop_generation_enabled: bool = False
single_clip_mode: bool = False
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
continuation_state: dict[str, Any] | None = None
"""Optional ``{kind, payload}`` dict; hydrated into
:class:`fastvideo.api.ContinuationState` server-side."""
class SegmentPromptSource(BaseModel):
"""Request a new segment using a specific prompt."""
type: Literal["segment_prompt_source"]
prompt: str
negative_prompt: str | None = None
source: Literal["curated", "enhanced", "user", "auto_extension"] = "user"
seed: int | None = None
num_inference_steps: int | None = None
guidance_scale: float | None = None
class SeedPromptsUpdated(BaseModel):
type: Literal["seed_prompts_updated"]
seed_prompts: list[str] = Field(default_factory=list)
class EnhancementUpdated(BaseModel):
type: Literal["enhancement_updated"]
enabled: bool
class AutoExtensionUpdated(BaseModel):
type: Literal["auto_extension_updated"]
enabled: bool
class LoopGenerationUpdated(BaseModel):
type: Literal["loop_generation_updated"]
enabled: bool
class GenerationPausedUpdated(BaseModel):
type: Literal["generation_paused_updated"]
paused: bool
class SnapshotState(BaseModel):
"""Request the current ``ContinuationState`` for export."""
type: Literal["snapshot_state"]
ClientMessage = Annotated[
Union[ # noqa: UP007 - Annotated requires Union for discriminator
SessionInitV2,
SegmentPromptSource,
SeedPromptsUpdated,
EnhancementUpdated,
AutoExtensionUpdated,
LoopGenerationUpdated,
GenerationPausedUpdated,
SnapshotState,
],
Field(discriminator="type"),
]
# ---------------------------------------------------------------------------
# Server → client
# ---------------------------------------------------------------------------
class QueueStatus(BaseModel):
type: Literal["queue_status"] = "queue_status"
position: int
queue_depth: int
class GpuAssigned(BaseModel):
type: Literal["gpu_assigned"] = "gpu_assigned"
gpu_id: int
session_timeout: int
class Ltx2StreamStart(BaseModel):
type: Literal["ltx2_stream_start"] = "ltx2_stream_start"
preset: str | None = None
width: int
height: int
fps: int
num_frames: int
class Ltx2SegmentStart(BaseModel):
type: Literal["ltx2_segment_start"] = "ltx2_segment_start"
segment_idx: int
prompt: str
total_steps: int
class StepComplete(BaseModel):
type: Literal["step_complete"] = "step_complete"
segment_idx: int
step: int
total_steps: int
stage: str = "denoise"
class MediaInit(BaseModel):
"""Descriptor for the fMP4 initialization segment that follows."""
type: Literal["media_init"] = "media_init"
segment_idx: int
mime: str = "video/mp4; codecs=\"avc1.64001f, mp4a.40.2\""
stream_id: str
mode: Literal["av_fmp4"] = "av_fmp4"
class MediaSegmentComplete(BaseModel):
type: Literal["media_segment_complete"] = "media_segment_complete"
segment_idx: int
stream_id: str
chunks: int
duration_ms: float | None = None
pts_base_ms: float | None = None
class Ltx2SegmentComplete(BaseModel):
type: Literal["ltx2_segment_complete"] = "ltx2_segment_complete"
segment_idx: int
generation_time_ms: float
e2e_latency_ms: float | None = None
class Ltx2StreamComplete(BaseModel):
type: Literal["ltx2_stream_complete"] = "ltx2_stream_complete"
reason: Literal["segment_cap", "stop_requested", "error"] = "stop_requested"
class SessionTimeout(BaseModel):
type: Literal["session_timeout"] = "session_timeout"
timeout_seconds: int
class ContinuationStateSnapshot(BaseModel):
type: Literal["continuation_state_snapshot"] = "continuation_state_snapshot"
state: dict[str, Any]
"""``{kind, payload}`` dict matching
:class:`fastvideo.api.ContinuationState`."""
class ErrorMessage(BaseModel):
type: Literal["error"] = "error"
code: Literal[
"session_rejected",
"invalid_message",
"preset_mismatch",
"gpu_unavailable",
"worker_failed",
"upstream_timeout",
"internal_error",
] = "internal_error"
message: str
retryable: bool = False
ServerMessage = Union[ # noqa: UP007 - pydantic Union handling
QueueStatus,
GpuAssigned,
Ltx2StreamStart,
Ltx2SegmentStart,
StepComplete,
MediaInit,
MediaSegmentComplete,
Ltx2SegmentComplete,
Ltx2StreamComplete,
SessionTimeout,
ContinuationStateSnapshot,
ErrorMessage,
]
def parse_client_message(raw: dict[str, Any]) -> ClientMessage:
"""Parse an incoming WebSocket dict into a typed client message.
Unknown ``type`` values raise :class:`pydantic.ValidationError`; the
server handler turns that into an ``error`` frame with
``code="invalid_message"``.
"""
from pydantic import TypeAdapter
return TypeAdapter(ClientMessage).validate_python(raw)
__all__ = [
"AutoExtensionUpdated",
"ClientMessage",
"ContinuationStateSnapshot",
"EnhancementUpdated",
"ErrorMessage",
"GenerationPausedUpdated",
"GpuAssigned",
"Ltx2SegmentComplete",
"Ltx2SegmentStart",
"Ltx2StreamComplete",
"Ltx2StreamStart",
"LoopGenerationUpdated",
"MediaInit",
"MediaSegmentComplete",
"QueueStatus",
"SeedPromptsUpdated",
"SegmentPromptSource",
"ServerMessage",
"SessionInitV2",
"SessionTimeout",
"SnapshotState",
"StepComplete",
"parse_client_message",
]
+520 -4
View File
@@ -1,15 +1,531 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-generator FastAPI + WebSocket streaming server."""
from __future__ import annotations
from fastvideo.api.schema import ServeConfig
import asyncio
import contextlib
import os
import time
from dataclasses import dataclass
from typing import Any, Protocol
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
InputConfig,
OutputConfig,
SamplingConfig,
ServeConfig,
)
from fastvideo.entrypoints.streaming.protocol import (
AutoExtensionUpdated,
ContinuationStateSnapshot,
EnhancementUpdated,
ErrorMessage,
GenerationPausedUpdated,
GpuAssigned,
LoopGenerationUpdated,
Ltx2SegmentComplete,
Ltx2SegmentStart,
Ltx2StreamComplete,
Ltx2StreamStart,
MediaInit,
MediaSegmentComplete,
QueueStatus,
SeedPromptsUpdated,
SegmentPromptSource,
SessionInitV2,
SnapshotState,
StepComplete,
parse_client_message,
)
from fastvideo.entrypoints.streaming.session import (
InvalidSessionTransition,
Session,
SessionManager,
SessionRejected,
SessionState,
)
from fastvideo.entrypoints.streaming.session_init_image import (
persist_session_init_image, )
from fastvideo.entrypoints.streaming.session_store import (
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# RFC 6455 WebSocket close codes used by the server.
_WS_CLOSE_UNSUPPORTED_DATA = 1003
_WS_CLOSE_TRY_AGAIN_LATER = 1013
def run_server(serve_config: ServeConfig) -> None:
"""Launch the streaming (WebSocket / Dynamo) server."""
class _GeneratorProto(Protocol):
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
def generate(self, request: GenerationRequest) -> Any:
...
@dataclass
class ServerState:
serve_config: ServeConfig
generator: _GeneratorProto
sessions: SessionManager
session_store: SessionStore
def build_app(
serve_config: ServeConfig,
generator: _GeneratorProto,
*,
session_store: SessionStore | None = None,
) -> FastAPI:
"""Build the FastAPI app used by :func:`run_server`.
Exposed so tests can drive the WebSocket endpoint in-process via
``starlette.testclient.TestClient(app).websocket_connect(...)``.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
"server; got None. Add a `streaming:` block to your serve config.")
sessions = SessionManager(
segment_cap=serve_config.streaming.generation_segment_cap,
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
)
state = ServerState(
serve_config=serve_config,
generator=generator,
sessions=sessions,
session_store=session_store or InMemorySessionStore(),
)
app = FastAPI(title="FastVideo Streaming")
@app.get("/health")
async def _health() -> JSONResponse:
return JSONResponse({
"status": "ok",
"sessions": len(state.sessions),
"stream_mode": state.serve_config.streaming.stream_mode,
})
@app.websocket("/v1/stream")
async def _stream(websocket: WebSocket) -> None:
await websocket.accept()
try:
session = state.sessions.create()
except SessionRejected as exc:
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
return
try:
await _handle_session(websocket, session, state)
except WebSocketDisconnect:
logger.info("session %s: client disconnected", session.id[:8])
except Exception: # pragma: no cover - defensive catch-all
logger.exception("session %s: unhandled error", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
finally:
_cleanup_session(session, state)
app.state.server_state = state
return app
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
"""Launch the streaming server.
Boots a :class:`fastvideo.VideoGenerator` from
``serve_config.generator`` unless ``generator`` is provided, then
serves ``build_app(...)`` via uvicorn.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
"got None. Add a `streaming:` block to your serve config.")
raise NotImplementedError("streaming server is not implemented yet")
import uvicorn
if generator is None:
from fastvideo import VideoGenerator # lazy to avoid boot cost
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
app = build_app(serve_config, generator)
uvicorn.run(
app,
host=serve_config.server.host,
port=serve_config.server.port,
)
async def _handle_session(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
init = await _read_init_message(websocket, session, state)
if init is None:
return
await _apply_session_init(session, init, state)
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
session.transition(SessionState.GPU_BINDING)
await _send_json(websocket, GpuAssigned(
gpu_id=0,
session_timeout=state.sessions.session_timeout_seconds,
))
session.transition(SessionState.ACTIVE)
await _send_json(websocket, _build_stream_start(session, state))
try:
await _run_segment_loop(websocket, session, state)
finally:
with contextlib.suppress(RuntimeError):
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
async def _read_init_message(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> SessionInitV2 | None:
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: init timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return None
except WebSocketDisconnect:
return None
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
return None
if not isinstance(parsed, SessionInitV2):
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
return None
return parsed
async def _reject_init(
websocket: WebSocket,
session: Session,
message: str,
close_reason: str,
) -> None:
await _send_error(websocket, "invalid_message", message, retryable=False)
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.REJECTED)
async def _apply_session_init(
session: Session,
init: SessionInitV2,
state: ServerState,
) -> None:
session.client_id = init.client_id
session.preset = init.preset
session.preset_label = init.preset_label
session.curated_prompts = list(init.curated_prompts)
session.enhancement_enabled = init.enhancement_enabled
session.auto_extension_enabled = init.auto_extension_enabled
session.loop_generation_enabled = init.loop_generation_enabled
session.single_clip_mode = init.single_clip_mode
session.stream_mode = init.stream_mode
if init.initial_image is not None:
# Decode + disk write off the event loop; payload is up to 32 MiB.
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
if image is not None:
session.metadata["session_init_image"] = image.path
if init.continuation_state is not None:
session.continuation_state = _coerce_state(init.continuation_state)
if session.continuation_state is not None:
state.session_store.store(session.id, session.continuation_state)
async def _run_segment_loop(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
cap = state.sessions.segment_cap
while True:
if session.segment_cap_reached(cap):
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
return
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: idle timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return
except WebSocketDisconnect:
return
session.touch()
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
continue
if isinstance(parsed, SnapshotState):
snap = state.session_store.snapshot(session.id)
if snap is None:
await _send_error(websocket,
"internal_error",
"no continuation state available for session",
retryable=False)
continue
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
continue
if isinstance(parsed, SegmentPromptSource):
await _run_segment(websocket, session, state, parsed)
continue
# Silently ignore unknown-but-valid types (additive-evolution
# rule in streaming.md).
_apply_toggle(session, parsed)
async def _run_segment(
websocket: WebSocket,
session: Session,
state: ServerState,
message: SegmentPromptSource,
) -> None:
request = _build_generation_request(session, message, state)
segment_idx = session.segment_idx
await _send_json(
websocket,
Ltx2SegmentStart(
segment_idx=segment_idx,
prompt=message.prompt,
total_steps=request.sampling.num_inference_steps,
))
start = time.perf_counter()
loop = asyncio.get_running_loop()
# TODO: executor-wrapped generate() cannot be cancelled, so a
# client disconnect mid-segment leaves the GPU work running to
# completion. Real cancellation needs the generate_async API.
try:
result = await loop.run_in_executor(None, state.generator.generate, request)
except Exception as exc:
logger.exception("session %s: generator failed", session.id[:8])
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
elapsed_ms = (time.perf_counter() - start) * 1000.0
frames = _extract_frames(result)
if not frames:
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
# Synchronous generator call has no per-step hook; emit one
# terminal StepComplete so observability wiring still sees the
# segment finish.
total = request.sampling.num_inference_steps
await _send_json(websocket, StepComplete(
segment_idx=segment_idx,
step=total,
total_steps=total,
stage="denoise",
))
encoder = FragmentedMP4Encoder(
width=request.sampling.width,
height=request.sampling.height,
fps=request.sampling.fps,
segment_idx=segment_idx,
)
chunks_relayed = 0
async with encoder:
init_sent = False
async for chunk in encoder.encode(frames):
if chunk.kind == "init":
await _send_json(websocket, MediaInit(
segment_idx=segment_idx,
stream_id=chunk.stream_id,
))
init_sent = True
await websocket.send_bytes(chunk.data)
if init_sent and chunk.kind == "media":
chunks_relayed += 1
await _send_json(
websocket,
MediaSegmentComplete(
segment_idx=segment_idx,
stream_id=encoder.stream_id,
chunks=chunks_relayed,
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
))
new_state = _extract_state(result)
if new_state is not None:
session.continuation_state = new_state
state.session_store.store(session.id, new_state)
session.segment_idx += 1
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ACTIVE)
await _send_json(
websocket,
Ltx2SegmentComplete(
segment_idx=segment_idx,
generation_time_ms=elapsed_ms,
e2e_latency_ms=elapsed_ms,
))
def _build_stream_start(
session: Session,
state: ServerState,
) -> Ltx2StreamStart:
default = state.serve_config.default_request
return Ltx2StreamStart(
preset=session.preset,
width=default.sampling.width,
height=default.sampling.height,
fps=default.sampling.fps,
num_frames=default.sampling.num_frames,
)
def _build_generation_request(
session: Session,
message: SegmentPromptSource,
state: ServerState,
) -> GenerationRequest:
# Start from the operator-pinned default_request to pick up the
# preset-selected sampling knobs; override with per-message values.
base = state.serve_config.default_request
sampling_kwargs: dict[str, Any] = {
"num_videos_per_prompt":
base.sampling.num_videos_per_prompt,
"seed":
message.seed if message.seed is not None else base.sampling.seed,
"num_frames":
base.sampling.num_frames,
"height":
base.sampling.height,
"width":
base.sampling.width,
"fps":
base.sampling.fps,
"num_inference_steps":
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
"guidance_scale":
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
}
request = GenerationRequest(
prompt=message.prompt,
negative_prompt=message.negative_prompt or base.negative_prompt,
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
sampling=SamplingConfig(**sampling_kwargs),
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
state=session.continuation_state,
)
return request
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
kind = raw.get("kind")
payload = raw.get("payload")
if not isinstance(kind, str) or not isinstance(payload, dict):
return None
return ContinuationState(kind=kind, payload=payload)
def _apply_toggle(session: Session, message: Any) -> None:
if isinstance(message, EnhancementUpdated):
session.enhancement_enabled = message.enabled
elif isinstance(message, AutoExtensionUpdated):
session.auto_extension_enabled = message.enabled
elif isinstance(message, LoopGenerationUpdated):
session.loop_generation_enabled = message.enabled
elif isinstance(message, GenerationPausedUpdated):
session.generation_paused = message.paused
elif isinstance(message, SeedPromptsUpdated):
session.curated_prompts = list(message.seed_prompts)
def _extract_frames(result: Any) -> list:
if hasattr(result, "frames"):
return list(result.frames or [])
if isinstance(result, dict):
return list(result.get("frames") or [])
return []
def _extract_state(result: Any) -> ContinuationState | None:
state = getattr(result, "state", None)
if state is None and isinstance(result, dict):
state = result.get("state")
if isinstance(state, ContinuationState):
return state
if isinstance(state, dict):
return _coerce_state(state)
return None
async def _send_json(websocket: WebSocket, message: Any) -> None:
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
await websocket.send_json(payload)
async def _send_error(
websocket: WebSocket,
code: str,
message: str,
*,
retryable: bool,
) -> None:
await _send_json(
websocket,
ErrorMessage(code=code, message=message, retryable=retryable),
)
def _cleanup_session(session: Session, state: ServerState) -> None:
state.sessions.close(session.id)
state.session_store.drop(session.id)
init_image_path = session.metadata.get("session_init_image")
if isinstance(init_image_path, str):
with contextlib.suppress(FileNotFoundError):
os.unlink(init_image_path)
__all__ = [
"ServerState",
"build_app",
"run_server",
]
+214
View File
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-connection session lifecycle for the streaming server.
Each WebSocket opens exactly one :class:`Session`. :class:`SessionManager`
enforces the ``generation_segment_cap`` and ``session_timeout_seconds``
budgets from :class:`fastvideo.api.StreamingConfig`.
"""
from __future__ import annotations
import enum
import time
import uuid
from dataclasses import dataclass, field
from typing import Any
from fastvideo.api.schema import ContinuationState
class SessionState(enum.Enum):
"""State-machine positions for a streaming session.
Transitions are server-owned. See
``docs/design/server_contracts/streaming.md`` for the full diagram.
"""
INITIALIZING = "initializing"
QUEUED = "queued"
GPU_BINDING = "gpu_binding"
ACTIVE = "active"
COMPLETE = "complete"
ERROR = "error"
TIMEOUT = "timeout"
REJECTED = "rejected"
_VALID_TRANSITIONS: dict[SessionState, frozenset[SessionState]] = {
SessionState.INITIALIZING:
frozenset({
SessionState.QUEUED,
SessionState.GPU_BINDING,
SessionState.REJECTED,
SessionState.ERROR,
}),
SessionState.QUEUED:
frozenset({
SessionState.GPU_BINDING,
SessionState.ERROR,
SessionState.TIMEOUT,
SessionState.REJECTED,
}),
SessionState.GPU_BINDING:
frozenset({
SessionState.ACTIVE,
SessionState.ERROR,
SessionState.TIMEOUT,
}),
SessionState.ACTIVE:
frozenset({
SessionState.ACTIVE,
SessionState.COMPLETE,
SessionState.ERROR,
SessionState.TIMEOUT,
}),
SessionState.COMPLETE:
frozenset(),
SessionState.ERROR:
frozenset(),
SessionState.TIMEOUT:
frozenset(),
SessionState.REJECTED:
frozenset(),
}
class InvalidSessionTransition(RuntimeError):
"""Raised when a session is asked to transition along an illegal edge."""
@dataclass
class Session:
id: str = field(default_factory=lambda: uuid.uuid4().hex)
state: SessionState = SessionState.INITIALIZING
created_at: float = field(default_factory=time.monotonic)
last_activity: float = field(default_factory=time.monotonic)
client_id: str | None = None
preset: str | None = None
preset_label: str | None = None
curated_prompts: list[str] = field(default_factory=list)
segment_idx: int = 0
enhancement_enabled: bool = False
auto_extension_enabled: bool = False
loop_generation_enabled: bool = False
single_clip_mode: bool = False
generation_paused: bool = False
stream_mode: str = "av_fmp4"
gpu_id: int | None = None
continuation_state: ContinuationState | None = None
metadata: dict[str, Any] = field(default_factory=dict)
def transition(self, target: SessionState) -> None:
"""Move to ``target`` if the edge is allowed.
Raises :class:`InvalidSessionTransition` on illegal moves. The
self-loop on ``ACTIVE`` is legal so the server can re-assert
ACTIVE on segment completion without special casing.
"""
allowed = _VALID_TRANSITIONS.get(self.state, frozenset())
if target not in allowed and target is not self.state:
raise InvalidSessionTransition(f"{self.state.value} -> {target.value} is not a valid "
f"session transition")
self.state = target
self.last_activity = time.monotonic()
def touch(self) -> None:
self.last_activity = time.monotonic()
def is_active(self) -> bool:
return self.state is SessionState.ACTIVE
def segment_cap_reached(self, cap: int) -> bool:
return self.segment_idx >= cap
class SessionManager:
"""Registers sessions and enforces per-server session limits."""
def __init__(
self,
*,
segment_cap: int,
session_timeout_seconds: int,
max_sessions: int = 1,
) -> None:
self._segment_cap = segment_cap
self._session_timeout_seconds = session_timeout_seconds
self._max_sessions = max_sessions
self._sessions: dict[str, Session] = {}
@property
def segment_cap(self) -> int:
return self._segment_cap
@property
def session_timeout_seconds(self) -> int:
return self._session_timeout_seconds
def create(self) -> Session:
if len(self._sessions) >= self._max_sessions:
raise SessionRejected(f"max sessions reached ({self._max_sessions})")
session = Session()
self._sessions[session.id] = session
return session
def get(self, session_id: str) -> Session | None:
return self._sessions.get(session_id)
def close(self, session_id: str) -> None:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
return session_id in self._sessions
def __len__(self) -> int:
return len(self._sessions)
def active_sessions(self) -> list[Session]:
return [s for s in self._sessions.values() if s.is_active()]
def reap_timed_out(self, now: float | None = None) -> list[str]:
"""Return the ids of sessions that have exceeded the idle timeout.
The caller is responsible for actually closing them — this
method only *identifies* dead sessions so the server can emit
``session_timeout`` frames before dropping the WebSocket.
TODO: unused until a background driver calls it. Per-connection
idle enforcement currently happens via asyncio.wait_for on
receive_json; this helper catches sessions stuck before any
receive (e.g. future QUEUED state) and is expected to be wired
into the GPU-pool reaper.
"""
now = now if now is not None else time.monotonic()
dead: list[str] = []
for sid, session in self._sessions.items():
if session.state in {
SessionState.COMPLETE,
SessionState.ERROR,
SessionState.TIMEOUT,
SessionState.REJECTED,
}:
continue
if now - session.last_activity > self._session_timeout_seconds:
dead.append(sid)
return dead
class SessionRejected(RuntimeError):
"""Raised when session creation fails (queue full, auth, etc.)."""
__all__ = [
"InvalidSessionTransition",
"Session",
"SessionManager",
"SessionRejected",
"SessionState",
]
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
"""Persist the initial-image blob attached to a streaming session."""
from __future__ import annotations
import base64
import binascii
import contextlib
import os
import tempfile
from dataclasses import dataclass
from typing import Any
_ACCEPTED_MIMES = {
"image/png": ".png",
"image/jpeg": ".jpg",
"image/jpg": ".jpg",
"image/webp": ".webp",
}
_MAX_IMAGE_BYTES = 32 * 1024 * 1024 # 32 MiB cap
@dataclass(frozen=True)
class SessionInitImage:
"""Location of the persisted init image.
Callers pass ``path`` to ``InputConfig.image_path``; ``display_name``
is only used for logs.
"""
path: str
display_name: str
mime: str
def persist_session_init_image(
payload: Any,
*,
output_dir: str | None = None,
) -> SessionInitImage | None:
"""Decode a client init-image blob and persist it to disk.
``payload`` shape (matches the internal UI protocol)::
{
"mime": "image/png",
"name": "ref.png",
"data": "<base64 bytes>",
}
Returns ``None`` when ``payload`` is falsy (no init image). Raises
:class:`ValueError` on schema / size / decode errors so the caller
can surface a user-facing ``error`` frame.
"""
if not payload:
return None
if not isinstance(payload, dict):
raise ValueError("session init image must be an object")
mime = payload.get("mime")
if mime not in _ACCEPTED_MIMES:
raise ValueError(f"session init image mime {mime!r} is not one of "
f"{sorted(_ACCEPTED_MIMES)}")
data_b64 = payload.get("data")
if not isinstance(data_b64, str):
raise ValueError("session init image data must be a base64 string")
try:
data = base64.b64decode(data_b64, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"session init image data is not valid base64: {exc}") from exc
if len(data) > _MAX_IMAGE_BYTES:
raise ValueError(f"session init image is {len(data)} bytes; limit is "
f"{_MAX_IMAGE_BYTES}")
if len(data) == 0:
raise ValueError("session init image data is empty")
ext = _ACCEPTED_MIMES[mime]
display_name = _sanitize_display_name(payload.get("name")) or f"init{ext}"
fd, path = tempfile.mkstemp(prefix="fastvideo-init-", suffix=ext, dir=output_dir)
try:
with os.fdopen(fd, "wb") as f:
f.write(data)
except Exception:
with contextlib.suppress(FileNotFoundError):
os.unlink(path)
raise
return SessionInitImage(path=path, display_name=display_name, mime=mime)
def _sanitize_display_name(name: Any) -> str | None:
if not isinstance(name, str):
return None
name = name.strip()
if not name:
return None
# Strip any path components — we only keep the leaf for logging.
return os.path.basename(name)
__all__ = [
"SessionInitImage",
"persist_session_init_image",
]
@@ -0,0 +1,206 @@
# SPDX-License-Identifier: Apache-2.0
"""Session state store for the FastVideo streaming server.
The streaming server keeps continuation state (decoded frames + audio
latents from the previous segment) server-side so the client doesn't
re-upload multi-megabyte tensors each WebSocket message. Two operations
are needed:
* ``snapshot(session_id) -> ContinuationState`` — serialize the current
state so it can be exported (e.g. over HTTP) or migrated to a
different server.
* ``hydrate(state) -> session_id`` — load a previously serialized state
into a new session (for resume-after-disconnect flows).
The store is an ABC with an :class:`InMemorySessionStore` default; Redis
or other backends can drop in without touching the pipeline.
Large tensor payloads (video frames, audio latents) are kept out of the
JSON payload via an accompanying :class:`BlobStore`. Both stores share a
process today; they are separate types so that a future implementation
can put blobs on S3 while keeping session metadata in Redis.
"""
from __future__ import annotations
import threading
import uuid
from abc import ABC, abstractmethod
from collections.abc import Iterator
from dataclasses import dataclass
from fastvideo.api.schema import ContinuationState
class BlobStore(ABC):
"""Opaque byte-blob storage keyed by id.
A :class:`ContinuationState` payload can reference large tensors
stored in a :class:`BlobStore` rather than inlining them, so the
JSON payload stays small when the state travels over the wire.
"""
@abstractmethod
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
"""Store ``data`` and return a blob id for later retrieval."""
@abstractmethod
def get(self, blob_id: str) -> bytes:
"""Load a previously stored blob. Raises ``KeyError`` if absent."""
@abstractmethod
def drop(self, blob_id: str) -> None:
"""Remove a blob. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, blob_id: str) -> bool:
...
@dataclass(frozen=True)
class _BlobRecord:
data: bytes
mime: str
class InMemoryBlobStore(BlobStore):
"""Thread-safe in-memory :class:`BlobStore` for single-process servers.
No eviction policy — callers are responsible for calling
:meth:`drop` when a blob's owning state is replaced or a session
ends. A redis- or filesystem-backed :class:`BlobStore` should
replace this when the streaming server lands as a real service
(PR 7.5+).
"""
def __init__(self) -> None:
self._blobs: dict[str, _BlobRecord] = {}
self._lock = threading.Lock()
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
blob_id = uuid.uuid4().hex
with self._lock:
self._blobs[blob_id] = _BlobRecord(data=data, mime=mime)
return blob_id
def get(self, blob_id: str) -> bytes:
with self._lock:
record = self._blobs.get(blob_id)
if record is None:
raise KeyError(f"Unknown blob id: {blob_id}")
return record.data
def drop(self, blob_id: str) -> None:
with self._lock:
self._blobs.pop(blob_id, None)
def __contains__(self, blob_id: str) -> bool:
with self._lock:
return blob_id in self._blobs
def __len__(self) -> int:
with self._lock:
return len(self._blobs)
class SessionStore(ABC):
"""Keyed store for per-session continuation state.
Implementations own the session-id → state mapping. The streaming
server calls :meth:`store` after each segment and :meth:`snapshot`
when a client explicitly asks for an exportable state handle.
"""
@abstractmethod
def store(self, session_id: str, state: ContinuationState) -> None:
"""Persist ``state`` for ``session_id``, replacing any prior value."""
@abstractmethod
def snapshot(self, session_id: str) -> ContinuationState | None:
"""Return the current state for ``session_id`` (or ``None``)."""
@abstractmethod
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
"""Install ``state`` as the starting point for a session.
When ``session_id`` is ``None`` the store allocates a fresh id
(UUID4); when provided the store uses it verbatim, overwriting
any prior state at that id.
"""
@abstractmethod
def drop(self, session_id: str) -> None:
"""Forget a session. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, session_id: str) -> bool:
...
@abstractmethod
def __iter__(self) -> Iterator[str]:
...
class InMemorySessionStore(SessionStore):
"""Thread-safe in-memory :class:`SessionStore`.
Default implementation used by single-process deployments; a future
Redis-backed store can be dropped in without changes to the server.
No eviction / TTL / bounded capacity — sessions only leave via
:meth:`drop`. The live streaming server (PR 7.5+) is responsible
for bounding growth and for dropping any :class:`BlobStore` blobs
referenced by a state when that state is replaced or a session
ends; this class does not know about blobs.
"""
def __init__(self) -> None:
self._sessions: dict[str, ContinuationState] = {}
self._lock = threading.Lock()
def store(self, session_id: str, state: ContinuationState) -> None:
with self._lock:
self._sessions[session_id] = state
def snapshot(self, session_id: str) -> ContinuationState | None:
with self._lock:
return self._sessions.get(session_id)
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
sid = session_id or uuid.uuid4().hex
with self._lock:
self._sessions[sid] = state
return sid
def drop(self, session_id: str) -> None:
with self._lock:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
with self._lock:
return session_id in self._sessions
def __iter__(self) -> Iterator[str]:
with self._lock:
return iter(list(self._sessions))
def __len__(self) -> int:
with self._lock:
return len(self._sessions)
__all__ = [
"BlobStore",
"InMemoryBlobStore",
"InMemorySessionStore",
"SessionStore",
]
+213
View File
@@ -0,0 +1,213 @@
# SPDX-License-Identifier: Apache-2.0
"""fMP4 stream encoder used by the streaming server.
The client's Media Source Extensions player needs a continuous fMP4
byte stream: first an *initialization segment* (``ftyp`` + ``moov``),
then one or more *media segments* (``moof`` + ``mdat``). We pipe raw
RGB frames into an ffmpeg subprocess configured for fragmented output
via ``-movflags empty_moov+default_base_moof+frag_keyframe+faststart``
and stream the bytes back out.
"""
from __future__ import annotations
import asyncio
import contextlib
import subprocess
import uuid
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
if TYPE_CHECKING:
import numpy as np
@dataclass
class FragmentedMP4Chunk:
"""A single fMP4 byte chunk emitted by :class:`FragmentedMP4Encoder`.
``kind`` identifies whether the chunk is the init segment (must be
fed into the client's ``SourceBuffer`` first) or a media fragment.
"""
kind: Literal["init", "media"]
data: bytes
stream_id: str
segment_idx: int
class FragmentedMP4Encoder:
"""Stream RGB frames in, fMP4 chunks out.
One encoder covers one segment. The server creates a new encoder
per :class:`ltx2_segment_start`` boundary so each segment becomes
one media fragment the client can append independently.
Example::
encoder = FragmentedMP4Encoder(width=1024, height=576, fps=24,
segment_idx=0)
async with encoder:
async for chunk in encoder.encode(frames):
await websocket.send_bytes(chunk.data)
"""
def __init__(
self,
*,
width: int,
height: int,
fps: int,
segment_idx: int,
stream_id: str | None = None,
ffmpeg_path: str = "ffmpeg",
preset: str = "ultrafast",
pixel_format_out: str = "yuv420p",
extra_args: list[str] | None = None,
) -> None:
self.width = width
self.height = height
self.fps = fps
self.segment_idx = segment_idx
self.stream_id = stream_id or uuid.uuid4().hex
self._ffmpeg_path = ffmpeg_path
self._preset = preset
self._pixel_format_out = pixel_format_out
self._extra_args = list(extra_args or [])
self._proc: subprocess.Popen | None = None
self._init_emitted = False
async def __aenter__(self) -> FragmentedMP4Encoder:
self._spawn()
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
await self.close()
def _spawn(self) -> None:
args = [
self._ffmpeg_path,
"-hide_banner",
"-loglevel",
"error",
"-f",
"rawvideo",
"-pix_fmt",
"rgb24",
"-s",
f"{self.width}x{self.height}",
"-r",
str(self.fps),
"-i",
"-",
"-c:v",
"libx264",
"-preset",
self._preset,
"-tune",
"zerolatency",
"-pix_fmt",
self._pixel_format_out,
"-movflags",
"empty_moov+default_base_moof+frag_keyframe+faststart",
"-f",
"mp4",
*self._extra_args,
"-",
]
# stderr → DEVNULL: with -loglevel error on, the only thing
# stderr would carry is unsolicited warnings. Piping without a
# reader deadlocks ffmpeg once the pipe buffer (~64 KiB) fills.
self._proc = subprocess.Popen( # noqa: S603
args,
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL,
bufsize=0,
)
async def encode(
self,
frames: list[np.ndarray] | AsyncIterator[np.ndarray],
) -> AsyncIterator[FragmentedMP4Chunk]:
"""Feed frames into ffmpeg and yield fMP4 chunks as they appear."""
if self._proc is None:
self._spawn()
assert self._proc is not None and self._proc.stdin is not None
proc = self._proc
loop = asyncio.get_running_loop()
async def _writer() -> None:
try:
if hasattr(frames, "__aiter__"):
async for frame in frames: # type: ignore[union-attr]
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
else:
for frame in frames: # type: ignore[assignment]
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
finally:
with contextlib.suppress(BrokenPipeError):
proc.stdin.close()
writer_task = asyncio.create_task(_writer())
try:
reader = proc.stdout
assert reader is not None
# Read in reasonably-sized chunks; MSE tolerates any size
# but we don't want to starve the event loop.
chunk_size = 64 * 1024
while True:
data = await loop.run_in_executor(None, reader.read, chunk_size)
if not data:
break
kind: Literal["init", "media"] = "init" if not self._init_emitted else "media"
self._init_emitted = True
yield FragmentedMP4Chunk(
kind=kind,
data=bytes(data),
stream_id=self.stream_id,
segment_idx=self.segment_idx,
)
finally:
await writer_task
async def close(self) -> None:
if self._proc is None:
return
proc = self._proc
self._proc = None
try:
if proc.stdin and not proc.stdin.closed:
proc.stdin.close()
except BrokenPipeError:
pass
loop = asyncio.get_running_loop()
try:
await asyncio.wait_for(
loop.run_in_executor(None, proc.wait),
timeout=5.0,
)
except asyncio.TimeoutError:
proc.kill()
await loop.run_in_executor(None, proc.wait)
def _write_frame(stdin, frame: np.ndarray) -> None:
import numpy as np
if not isinstance(frame, np.ndarray):
raise TypeError("fMP4 encoder frames must be numpy.ndarray")
if frame.dtype != np.uint8:
frame = frame.astype(np.uint8)
if frame.ndim != 3 or frame.shape[-1] != 3:
raise ValueError("fMP4 encoder frames must be HxWx3 uint8 RGB; got "
f"shape={frame.shape}, dtype={frame.dtype}")
with contextlib.suppress(BrokenPipeError):
stdin.write(frame.tobytes())
__all__ = [
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
]
+78 -36
View File
@@ -627,18 +627,32 @@ class VideoGenerator:
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# 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())
# 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())
# Save output if requested
if batch.save_video:
if self._is_image_workload():
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():
# Image workloads (t2i, i2i, …): save the first frame as PNG.
imageio.imwrite(output_path, frames[0])
logger.info("Saved image to %s", output_path)
@@ -655,7 +669,11 @@ class VideoGenerator:
"prompts": prompt,
"samples": samples if batch.return_frames else None,
"frames": frames if batch.return_frames else None,
"audio": output_batch.extra.get("audio") 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"),
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
@@ -683,7 +701,55 @@ 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,
@@ -696,37 +762,13 @@ 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")
# 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())
num_channels = cls._write_pcm_wav(wav_path, audio, sample_rate)
layout = "stereo" if num_channels == 2 else "mono"
# Open input video and audio
input_video = av.open(video_path)
+11 -1
View File
@@ -844,6 +844,13 @@ class TrainingArgs(FastVideoArgs):
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
# CFG scale applied to the real (teacher) score in the DMD loss, using the
# parameterization `x = x_cond + w * (x_cond - x_uncond)`. This differs
# from the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)` by an
# offset of 1: `w_here = w_standard - 1`. So `w=0` recovers the
# conditional output, `w=-1` recovers the unconditional output, and the
# default 3.5 corresponds to a standard CFG scale of 4.5. Matches the
# original DMD2 reference implementation.
real_score_guidance_scale: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
@@ -1104,7 +1111,10 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--real-score-guidance-scale",
type=float,
default=TrainingArgs.real_score_guidance_scale,
help="Teacher guidance scale")
help=("Teacher CFG scale for the real score in the DMD loss. Uses "
"the parameterization x_cond + w * (x_cond - x_uncond), so "
"w=0 -> cond, w=-1 -> uncond, and the relation to standard "
"CFG is w_standard = w + 1 (default 3.5 == standard 4.5)."))
parser.add_argument("--fake-score-learning-rate",
type=float,
default=TrainingArgs.fake_score_learning_rate,
+35 -1
View File
@@ -556,7 +556,10 @@ class CosmosTransformer3DModel(BaseDiT):
self.extra_pos_embed_type = config.extra_pos_embed_type
# 1. Patch Embedding
patch_embed_in_channels = config.in_channels + 1 if config.concat_padding_mask else config.in_channels
# 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)
self.patch_embed = CosmosPatchEmbed(patch_embed_in_channels,
inner_dim,
config.patch_size,
@@ -617,6 +620,28 @@ 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)
@@ -626,6 +651,10 @@ 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
)
@@ -703,6 +732,11 @@ 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
@@ -0,0 +1,389 @@
# 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
@@ -0,0 +1,214 @@
# 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
@@ -1,224 +0,0 @@
# 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,6 +95,10 @@ 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:
@@ -998,6 +1002,61 @@ 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,6 +86,9 @@ _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
@@ -0,0 +1,376 @@
# 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
@@ -0,0 +1,122 @@
# 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,23 +22,18 @@ class Cosmos2VideoToWorldPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler", "safety_checker"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
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
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
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
@@ -0,0 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# Importing continuation registers the "ltx2.v1" continuation kind with
# the public compat layer so GenerationRequest.state(kind="ltx2.v1") is
# recognized on the public API boundary.
from fastvideo.pipelines.basic.ltx2 import continuation # noqa: F401
@@ -0,0 +1,386 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed continuation state for the LTX-2 streaming pipeline.
Segment N+1 conditions on segment N's trailing decoded frames and
denoised audio latents. The streaming runtime used to hold this state as
per-worker globals; lifting it into a typed, JSON-serializable object
lets clients snapshot, migrate, or round-trip it through an HTTP/RPC
boundary. The envelope ``ContinuationState(kind, payload)`` is the
shared public API; the typed class here owns the LTX-2 payload shape.
Serialization contract:
* Video frames → PNG bytes + base64, or a :class:`BlobStore` id.
* Audio latents → a self-describing safetensors blob + base64, or a
:class:`BlobStore` id. safetensors preserves ``bfloat16``, which a
raw-numpy round-trip cannot.
* The returned payload is always a plain JSON-serializable dict.
"""
from __future__ import annotations
import base64
from collections.abc import Mapping
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from fastvideo.api.compat import register_continuation_kind
from fastvideo.api.schema import ContinuationState
if TYPE_CHECKING:
import numpy as np
import torch
from fastvideo.entrypoints.streaming.session_store import BlobStore
LTX2_CONTINUATION_KIND = "ltx2.v1"
"""Public ``ContinuationState.kind`` for LTX-2 payloads."""
LTX2_CONTINUATION_SCHEMA_VERSION = 1
"""Payload schema version carried inside ``payload.schema_version``."""
DEFAULT_INLINE_THRESHOLD_BYTES = 2 * 1024 * 1024
"""Tensors larger than this go to the blob store (if available). 2 MiB
is below typical single-JSON-message limits (Dynamo: 4 MiB, Postgres
TOAST: 1 GiB) and well above per-frame PNG payloads (~200 KiB at
512x512)."""
@dataclass
class LTX2ContinuationState:
"""Typed LTX-2 continuation state carried between streaming segments.
``video_frames`` hold trailing decoded RGB frames (uint8 HxWx3) from
segment N for conditioning segment N+1 via the VAE encode path.
``audio_latents`` is the cached denoised audio latent tensor of shape
``[B, C, T, mel]`` that segment N+1 will copy into the overlap
region of its clean-latent conditioning.
Most fields map 1:1 onto the internal gpu_pool's per-worker state;
the only new concept is the ``*_blob_id`` fields, which allow large
tensors to live outside the JSON payload. See module docstring.
"""
segment_index: int = 0
"""Index of the *just-completed* segment. Segment 0 has no history;
state returned after segment 0 carries ``segment_index=0`` and the
caller uses ``segment_index + 1`` as the next segment number."""
video_frames: list[np.ndarray] | None = None
"""Trailing decoded frames, each an RGB uint8 ``np.ndarray`` shaped
``(H, W, 3)``. ``None`` when the state is blob-backed or unset."""
video_frames_blob_id: str | None = None
"""Blob store id when the frames live outside the payload."""
video_conditioning_frame_idx: int = 0
"""Target frame index inside the next segment that the trailing
frames align with (matches the LTX-2 ``ltx2_video_conditions``
tuple's ``frame_idx`` slot)."""
video_conditioning_strength: float = 1.0
"""Conditioning strength in [0, 1]. Matches the ``ltx2_video_
conditions`` tuple's strength slot."""
audio_latents: torch.Tensor | None = None
"""Denoised audio latent tensor of shape ``[B, C, T, mel]``.
``None`` when the state is blob-backed or unset."""
audio_latents_blob_id: str | None = None
"""Blob store id when audio latents live outside the payload."""
audio_sample_rate: int | None = None
"""Sample rate for the audio side (e.g. 24000)."""
audio_conditioning_num_frames: int = 0
"""Number of trailing audio frames that carry over as clean context
into segment N+1."""
audio_conditioning_strength: float = 1.0
"""Clean-latent mask value applied to the overlap region; 0.0 keeps
the cached audio entirely, 1.0 renoises from scratch."""
video_position_offset_sec: float = 0.0
"""Seconds by which video RoPE is shifted forward so the audio
prefix can sit at ``t >= 0`` when audio conditioning is longer than
video conditioning."""
metadata: dict[str, Any] = field(default_factory=dict)
"""Opaque metadata bag for forward-compat fields that don't need
their own typed slot yet (e.g. custom knob experiments)."""
def to_continuation_state(
self,
*,
blob_store: BlobStore | None = None,
inline_threshold_bytes: int = DEFAULT_INLINE_THRESHOLD_BYTES,
) -> ContinuationState:
"""Serialize into a public :class:`ContinuationState`.
When ``blob_store`` is given, tensors larger than
``inline_threshold_bytes`` are stored via
:meth:`BlobStore.put` and referenced by id; otherwise all data
is base64-encoded inline. The payload is always a plain
JSON-serializable dict.
"""
payload: dict[str, Any] = {
"schema_version": LTX2_CONTINUATION_SCHEMA_VERSION,
"segment_index": int(self.segment_index),
"video_conditioning_frame_idx": int(self.video_conditioning_frame_idx),
"video_conditioning_strength": float(self.video_conditioning_strength),
"audio_conditioning_num_frames": int(self.audio_conditioning_num_frames),
"audio_conditioning_strength": float(self.audio_conditioning_strength),
"video_position_offset_sec": float(self.video_position_offset_sec),
"metadata": dict(self.metadata),
}
if self.audio_sample_rate is not None:
payload["audio_sample_rate"] = int(self.audio_sample_rate)
video_payload = self._encode_video_frames(
blob_store=blob_store,
inline_threshold_bytes=inline_threshold_bytes,
)
if video_payload is not None:
payload["video"] = video_payload
audio_payload = self._encode_audio_latents(
blob_store=blob_store,
inline_threshold_bytes=inline_threshold_bytes,
)
if audio_payload is not None:
payload["audio"] = audio_payload
return ContinuationState(
kind=LTX2_CONTINUATION_KIND,
payload=payload,
)
@classmethod
def from_continuation_state(
cls,
state: ContinuationState,
*,
blob_store: BlobStore | None = None,
) -> LTX2ContinuationState:
"""Rebuild a typed state from a public :class:`ContinuationState`.
Raises :class:`ValueError` when the kind doesn't match or the
schema version is unsupported.
"""
if state.kind != LTX2_CONTINUATION_KIND:
raise ValueError(f"Expected ContinuationState.kind={LTX2_CONTINUATION_KIND!r}, "
f"got {state.kind!r}")
payload = state.payload or {}
version = int(payload.get("schema_version", LTX2_CONTINUATION_SCHEMA_VERSION))
if version != LTX2_CONTINUATION_SCHEMA_VERSION:
raise ValueError(f"Unsupported LTX-2 continuation schema_version={version}; "
f"this build expects {LTX2_CONTINUATION_SCHEMA_VERSION}")
out = cls(
segment_index=int(payload.get("segment_index", 0)),
video_conditioning_frame_idx=int(payload.get("video_conditioning_frame_idx", 0)),
video_conditioning_strength=float(payload.get("video_conditioning_strength", 1.0)),
audio_sample_rate=(int(payload["audio_sample_rate"]) if "audio_sample_rate" in payload else None),
audio_conditioning_num_frames=int(payload.get("audio_conditioning_num_frames", 0)),
audio_conditioning_strength=float(payload.get("audio_conditioning_strength", 1.0)),
video_position_offset_sec=float(payload.get("video_position_offset_sec", 0.0)),
metadata=dict(payload.get("metadata") or {}),
)
video = payload.get("video")
if isinstance(video, Mapping):
cls._decode_video_frames(out, video, blob_store=blob_store)
audio = payload.get("audio")
if isinstance(audio, Mapping):
cls._decode_audio_latents(out, audio, blob_store=blob_store)
return out
# ------------------------------------------------------------------
# Video frame helpers
# ------------------------------------------------------------------
def _encode_video_frames(
self,
*,
blob_store: BlobStore | None,
inline_threshold_bytes: int,
) -> dict[str, Any] | None:
if self.video_frames_blob_id is not None:
return {"blob_id": self.video_frames_blob_id}
if not self.video_frames:
return None
encoded = [_encode_png(frame) for frame in self.video_frames]
total = sum(len(b) for b in encoded)
if blob_store is not None and total > inline_threshold_bytes:
concatenated = _pack_frame_blobs(encoded)
blob_id = blob_store.put(
concatenated,
mime="application/x-fastvideo-frames+png",
)
return {"blob_id": blob_id, "frame_count": len(encoded)}
return {
"frames_b64": [base64.b64encode(b).decode("ascii") for b in encoded],
}
@staticmethod
def _decode_video_frames(
out: LTX2ContinuationState,
video: Mapping[str, Any],
*,
blob_store: BlobStore | None,
) -> None:
blob_id = video.get("blob_id")
if isinstance(blob_id, str):
if blob_store is None:
out.video_frames_blob_id = blob_id
return
raw = blob_store.get(blob_id)
encoded = _unpack_frame_blobs(raw)
out.video_frames = [_decode_png(b) for b in encoded]
return
frames_b64 = video.get("frames_b64")
if isinstance(frames_b64, list):
decoded = [_decode_png(base64.b64decode(b)) for b in frames_b64 if isinstance(b, str)]
out.video_frames = decoded or None
# ------------------------------------------------------------------
# Audio latent helpers
# ------------------------------------------------------------------
def _encode_audio_latents(
self,
*,
blob_store: BlobStore | None,
inline_threshold_bytes: int,
) -> dict[str, Any] | None:
if self.audio_latents_blob_id is not None:
return {"blob_id": self.audio_latents_blob_id}
if self.audio_latents is None:
return None
raw = _tensor_to_safetensors_bytes(self.audio_latents)
if blob_store is not None and len(raw) > inline_threshold_bytes:
blob_id = blob_store.put(
raw,
mime="application/x-fastvideo-tensor+safetensors",
)
return {"blob_id": blob_id}
return {"safetensors_b64": base64.b64encode(raw).decode("ascii")}
@staticmethod
def _decode_audio_latents(
out: LTX2ContinuationState,
audio: Mapping[str, Any],
*,
blob_store: BlobStore | None,
) -> None:
blob_id = audio.get("blob_id")
if isinstance(blob_id, str):
if blob_store is None:
out.audio_latents_blob_id = blob_id
return
raw = blob_store.get(blob_id)
out.audio_latents = _safetensors_bytes_to_tensor(raw)
return
data_b64 = audio.get("safetensors_b64")
if isinstance(data_b64, str):
out.audio_latents = _safetensors_bytes_to_tensor(base64.b64decode(data_b64))
def _encode_png(frame: np.ndarray) -> bytes:
"""Encode an ``(H, W, 3)`` uint8 RGB frame as PNG bytes."""
import numpy as np
from PIL import Image
if not isinstance(frame, np.ndarray):
raise TypeError(f"LTX2 continuation frame must be a numpy ndarray, got {type(frame).__name__}")
if frame.dtype != np.uint8 or frame.ndim != 3 or frame.shape[-1] != 3:
raise ValueError("LTX2 continuation frame must be uint8 HxWx3 RGB; got "
f"dtype={frame.dtype}, shape={frame.shape}")
import io
buffer = io.BytesIO()
Image.fromarray(frame).save(buffer, format="PNG")
return buffer.getvalue()
def _decode_png(data: bytes) -> np.ndarray:
import io
import numpy as np
from PIL import Image
img = Image.open(io.BytesIO(data)).convert("RGB")
return np.array(img, dtype=np.uint8)
def _pack_frame_blobs(encoded: list[bytes]) -> bytes:
"""Pack multiple PNG blobs into a single blob for blob-store storage.
Format: ``[4-byte big-endian count][4-byte len][png][4-byte len][png]...``.
"""
parts: list[bytes] = [len(encoded).to_bytes(4, "big")]
for blob in encoded:
parts.append(len(blob).to_bytes(4, "big"))
parts.append(blob)
return b"".join(parts)
def _unpack_frame_blobs(raw: bytes) -> list[bytes]:
if len(raw) < 4:
raise ValueError("frame blob truncated: missing count header")
count = int.from_bytes(raw[:4], "big")
# Each frame contributes at least a 4-byte length prefix, so a
# declared count larger than (len(raw) - 4) // 4 cannot fit and
# would otherwise cause an O(count) allocation loop on malformed
# input.
if count > (len(raw) - 4) // 4:
raise ValueError(f"frame blob declares {count} frames but buffer holds at most "
f"{(len(raw) - 4) // 4}")
out: list[bytes] = []
cursor = 4
for index in range(count):
if cursor + 4 > len(raw):
raise ValueError(f"frame blob truncated at frame {index} length header")
length = int.from_bytes(raw[cursor:cursor + 4], "big")
cursor += 4
if cursor + length > len(raw):
raise ValueError(f"frame blob truncated at frame {index} payload")
out.append(raw[cursor:cursor + length])
cursor += length
return out
def _tensor_to_safetensors_bytes(tensor: Any) -> bytes:
"""Serialize a torch tensor to a self-describing safetensors blob.
Uses the in-memory safetensors API so the wire format preserves
dtype (including ``bfloat16``, which a raw-numpy path cannot) and
shape without needing sidecar metadata.
"""
import torch
from safetensors.torch import save as st_save
if isinstance(tensor, torch.Tensor):
return st_save({"t": tensor.detach().cpu()})
import numpy as np
if isinstance(tensor, np.ndarray):
return st_save({"t": torch.from_numpy(np.ascontiguousarray(tensor))})
raise TypeError("LTX2 audio_latents must be a torch.Tensor or numpy.ndarray, got "
f"{type(tensor).__name__}")
def _safetensors_bytes_to_tensor(raw: bytes) -> Any:
from safetensors.torch import load as st_load
return st_load(raw)["t"]
register_continuation_kind(LTX2_CONTINUATION_KIND)
__all__ = [
"DEFAULT_INLINE_THRESHOLD_BYTES",
"LTX2ContinuationState",
"LTX2_CONTINUATION_KIND",
"LTX2_CONTINUATION_SCHEMA_VERSION",
]
+35 -1
View File
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX2 model family pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_stage_override_fields, )
_LTX2_NEGATIVE_PROMPT = ("blurry, out of focus, overexposed, underexposed, low contrast, "
"washed out colors, excessive noise, grainy texture, poor lighting, "
@@ -30,6 +32,13 @@ _DENOISE_STAGE = PresetStageSpec(
}),
)
_REFINE_STAGE = PresetStageSpec(
name="refine",
kind="refinement",
description="Latent-upsample + second-pass refine",
allowed_overrides=refine_stage_override_fields(),
)
LTX2_BASE = InferencePreset(
name="ltx2_base",
version=1,
@@ -77,4 +86,29 @@ LTX2_DISTILLED = InferencePreset(
},
)
ALL_PRESETS = (LTX2_BASE, LTX2_DISTILLED)
LTX2_TWO_STAGE = InferencePreset(
name="ltx2_two_stage",
version=1,
model_family="ltx2",
description="LTX-2 distilled with 2x spatial refine (stage 1 half-res + stage 2 upsample+denoise)",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, _REFINE_STAGE),
defaults={
"seed": 10,
"height": 1024,
"width": 1536,
"num_frames": 121,
"fps": 24,
"guidance_scale": 1.0,
"num_inference_steps": 8,
"negative_prompt": "",
},
stage_defaults={
"refine": {
"num_inference_steps": 2,
"guidance_scale": 1.0,
},
},
)
ALL_PRESETS = (LTX2_BASE, LTX2_DISTILLED, LTX2_TWO_STAGE)
@@ -0,0 +1,59 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed override surfaces for the LTX-2 two-stage refine flow.
* ``preset_overrides.refine`` — init-time knobs (see
:class:`LTX2RefinePresetOverride`).
* ``stage_overrides.refine`` — per-request knobs (see
:class:`LTX2RefineStageOverride`).
Asset paths live on :class:`~fastvideo.api.schema.ComponentConfig`
(``upsampler_weights`` and ``lora_path``).
"""
from __future__ import annotations
from dataclasses import asdict, dataclass, fields
from typing import Any
@dataclass
class LTX2RefinePresetOverride:
"""Init-time refine wiring under ``preset_overrides.refine``."""
enabled: bool | None = None
add_noise: bool | None = None
@dataclass
class LTX2RefineStageOverride:
"""Per-request refine tuning under ``stage_overrides.refine``."""
# Stage-2 refine only validates 2 (reduced) and 3 (official distilled)
# sigma schedules; other values raise at pipeline construction.
num_inference_steps: int | None = None
guidance_scale: float | None = None
image_crf: int | None = None
video_position_offset_sec: float | None = None
def refine_override_to_dict(override: LTX2RefinePresetOverride | LTX2RefineStageOverride, ) -> dict[str, Any]:
"""Serialise a refine override, dropping ``None`` entries so only
user-set fields reach ``preset_overrides.refine`` or
``stage_overrides.refine``."""
return {k: v for k, v in asdict(override).items() if v is not None}
def refine_preset_override_fields() -> frozenset[str]:
return frozenset(f.name for f in fields(LTX2RefinePresetOverride))
def refine_stage_override_fields() -> frozenset[str]:
return frozenset(f.name for f in fields(LTX2RefineStageOverride))
__all__ = [
"LTX2RefinePresetOverride",
"LTX2RefineStageOverride",
"refine_override_to_dict",
"refine_preset_override_fields",
"refine_stage_override_fields",
]
@@ -0,0 +1,17 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2 family pipeline stages."""
from fastvideo.pipelines.basic.ltx2.stages.ltx2_audio_decoding import (
LTX2AudioDecodingStage, )
from fastvideo.pipelines.basic.ltx2.stages.ltx2_denoising import (
LTX2DenoisingStage, )
from fastvideo.pipelines.basic.ltx2.stages.ltx2_latent_preparation import (
LTX2LatentPreparationStage, )
from fastvideo.pipelines.basic.ltx2.stages.ltx2_text_encoding import (
LTX2TextEncodingStage, )
__all__ = [
"LTX2AudioDecodingStage",
"LTX2DenoisingStage",
"LTX2LatentPreparationStage",
"LTX2TextEncodingStage",
]
@@ -26,7 +26,7 @@ MATRIXGAME_I2V = InferencePreset(
"fps": 25,
"guidance_scale": 1.0,
"num_inference_steps": 3,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,63 @@
# 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)
@@ -0,0 +1,125 @@
# 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
@@ -0,0 +1,12 @@
# 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",
]
@@ -0,0 +1,98 @@
# 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
@@ -0,0 +1,67 @@
# 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
@@ -0,0 +1,207 @@
# 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
@@ -0,0 +1,174 @@
# 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": None,
"negative_prompt": "",
},
)
@@ -44,7 +44,7 @@ TURBO_T2V_14B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -62,7 +62,7 @@ TURBO_I2V_A14B = InferencePreset(
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": None,
"negative_prompt": "",
},
)
@@ -17,6 +17,8 @@ import torch
if TYPE_CHECKING:
from torchcodec.decoders import VideoDecoder
from fastvideo.api.schema import ContinuationState
import time
from collections import OrderedDict
@@ -190,6 +192,21 @@ 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
@@ -206,6 +223,9 @@ class ForwardBatch:
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None
continuation_state: "ContinuationState | None" = None
return_continuation_state: bool = False
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
@@ -0,0 +1,193 @@
# 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()
@@ -0,0 +1,196 @@
# 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()
+6 -4
View File
@@ -25,10 +25,12 @@ from fastvideo.pipelines.stages.latent_preparation import (Cosmos25LatentPrepara
Cosmos25AutoLatentPreparationStage,
Cosmos25T2WLatentPreparationStage,
Cosmos25V2WLatentPreparationStage, LatentPreparationStage)
from fastvideo.pipelines.stages.ltx2_audio_decoding import LTX2AudioDecodingStage
from fastvideo.pipelines.stages.ltx2_denoising import LTX2DenoisingStage
from fastvideo.pipelines.stages.ltx2_latent_preparation import (LTX2LatentPreparationStage)
from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
from fastvideo.pipelines.basic.ltx2.stages import (
LTX2AudioDecodingStage,
LTX2DenoisingStage,
LTX2LatentPreparationStage,
LTX2TextEncodingStage,
)
from fastvideo.pipelines.stages.matrixgame_denoising import (MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
+189 -166
View File
@@ -518,13 +518,42 @@ class DenoisingStage(PipelineStage):
class CosmosDenoisingStage(DenoisingStage):
"""
Denoising stage for Cosmos models using FlowMatchEulerDiscreteScheduler.
"""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.
"""
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,
@@ -533,199 +562,188 @@ 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
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
if hasattr(self.transformer, 'module'):
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_max = 80.0
sigma_min = 0.002
sigma_data = 1.0
final_sigmas_type = "sigma_min"
sigma_data = float(getattr(self.scheduler.config, "sigma_data", 1.0))
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)
self.scheduler.set_timesteps(
num_inference_steps,
device=latents.device,
)
timesteps = self.scheduler.timesteps
if (hasattr(self.scheduler.config, 'final_sigmas_type')
# Clamp terminal sigma to sigma_min (avoid zero).
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)
unconditioning_latents = conditioning_latents
conditioning_latents = getattr(
batch,
"conditioning_latents",
None,
)
cond_indicator = getattr(batch, "cond_indicator", None)
uncond_indicator = getattr(
batch,
"uncond_indicator",
None,
)
with self.progress_bar(total=num_inference_steps) as progress_bar:
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:
for i, t in enumerate(timesteps):
if hasattr(self, 'interrupt') and self.interrupt:
if hasattr(self, "interrupt") and self.interrupt:
continue
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
sigma = self.scheduler.sigmas[i]
is_aug_greater = bool(augment_sigma >= sigma)
timestep = current_t.view(1, 1, 1, 1, 1).expand(latents.size(0), -1, latents.size(2), -1,
-1) # [B, 1, T, 1, 1]
# 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)
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
# The model expects timestep = sigma * 1000
# (FlowMatchEulerDiscreteScheduler convention).
timestep_expanded = t.expand(latents.shape[0], ).to(target_dtype)
cond_latent = latents * c_in
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)
if hasattr(
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,
batch,
'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
))
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))
else:
logger.warning(
"Step %s: Missing conditioning data - cond_indicator: %s, conditioning_latents: %s", i,
hasattr(batch, 'cond_indicator'), conditioning_latents is not None)
final_x0 = cond_x0
cond_latent = cond_latent.to(target_dtype)
# Convert x0 to velocity for
# FlowMatchEulerDiscreteScheduler.
velocity = (latents - final_x0) / sigma.clamp(min=1e-6)
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]
latents = self.scheduler.step(
velocity,
t,
latents,
return_dict=False,
)[0]
progress_bar.update()
batch.latents = latents
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
@@ -771,16 +789,21 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
},
)
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
# 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
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,6 +524,22 @@ class ImageVAEEncodingStage(PipelineStage):
image = resize(image, height, width, resize_mode=resize_mode)
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
elif isinstance(image, torch.Tensor):
# VideoTransformStage delivers uint8 [0, 255] frames via batch.pil_image
# for the I2V preprocessing path. Convert here (not at the source) because
# batch.pil_image is also consumed as uint8 by ImageEncodingStage (HF
# processor does its own rescale) and by record_schema.py for parquet.
if image.dtype == torch.uint8:
image = image.float() / 255.0
elif not image.dtype.is_floating_point:
raise ValueError(f"preprocess() expected uint8 or float tensor, got {image.dtype}")
image_min = image.min()
image_max = image.max()
if image_max > 1.0 + 1e-4 or image_min < -1.0 - 1e-4:
raise ValueError("preprocess() expected tensor in [0, 1] or [-1, 1], got "
f"range [{image_min.item():.3f}, {image_max.item():.3f}]")
else:
raise TypeError(f"preprocess() expected PIL.Image or torch.Tensor, got {type(image)}")
do_normalize = True
if image.min() < 0:
+46 -1
View File
@@ -26,7 +26,7 @@ from fastvideo.configs.pipelines.hunyuan15 import (Hunyuan15T2V480PConfig, Hunyu
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionI2V_A14B_Config,
TurboDiffusionT2V_14B_Config,
@@ -48,6 +48,7 @@ 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
@@ -242,6 +243,47 @@ 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,
@@ -786,6 +828,8 @@ 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 (
@@ -803,6 +847,7 @@ def _register_presets() -> None:
LTX2_PRESETS,
MATRIXGAME_PRESETS,
SD35_PRESETS,
STABLE_AUDIO_PRESETS,
TURBODIFFUSION_PRESETS,
WAN_PRESETS,
)
+16 -4
View File
@@ -519,7 +519,7 @@ def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch)
cli_main.main()
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, monkeypatch):
config_path = tmp_path / "serve-streaming.yaml"
config_path.write_text(
"generator:\n"
@@ -530,9 +530,21 @@ def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
)
args, _ = _parse_serve_args(["--config", str(config_path)])
with pytest.raises(NotImplementedError,
match="streaming server is not implemented"):
ServeSubcommand().cmd(args)
captured: dict[str, object] = {}
def fake_run_server(serve_config, *, generator=None):
captured["serve_config"] = serve_config
def fail_if_called(*_args, **_kwargs):
raise AssertionError("OpenAI server must not run when streaming is set")
monkeypatch.setattr(streaming_server, "run_server", fake_run_server)
monkeypatch.setattr(api_server, "run_server", fail_if_called)
ServeSubcommand().cmd(args)
serve_config = captured["serve_config"]
assert serve_config.streaming is not None
assert serve_config.streaming.stream_mode == "av_fmp4"
def test_streaming_run_server_rejects_missing_streaming_block():
@@ -0,0 +1,226 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for ``fastvideo.api.compat`` translation helpers covering the
typed CompileConfig + PipelineSelection.vae_tiling surfaces promoted in
PR 6.
"""
from __future__ import annotations
from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
legacy_from_pretrained_to_config,
)
from fastvideo.api.schema import CompileConfig, GeneratorConfig
class TestLegacyTorchCompileKwargsTranslation:
"""Legacy ``torch_compile_kwargs={...}`` gets split across the four
first-class :class:`CompileConfig` fields and anything unknown falls
into ``extras``."""
def test_all_typed_keys_promoted(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{
"enable_torch_compile": True,
"torch_compile_kwargs": {
"backend": "inductor",
"fullgraph": True,
"mode": "max-autotune-no-cudagraphs",
"dynamic": False,
},
},
)
compile_config = config.engine.compile
assert compile_config.enabled is True
assert compile_config.backend == "inductor"
assert compile_config.fullgraph is True
assert compile_config.mode == "max-autotune-no-cudagraphs"
assert compile_config.dynamic is False
assert compile_config.extras == {}
def test_unknown_keys_land_in_extras(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{
"enable_torch_compile": True,
"torch_compile_kwargs": {
"backend": "inductor",
"options": {"triton.cudagraphs": False},
"disable": False,
},
},
)
compile_config = config.engine.compile
assert compile_config.backend == "inductor"
assert compile_config.extras == {
"options": {"triton.cudagraphs": False},
"disable": False,
}
def test_empty_kwargs_produces_empty_extras(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"torch_compile_kwargs": {}},
)
compile_config = config.engine.compile
assert compile_config.extras == {}
assert compile_config.backend is None
class TestCompileConfigRoundTrip:
"""typed CompileConfig -> FastVideoArgs.torch_compile_kwargs
reconstruction drops ``None`` typed fields and merges ``extras``."""
def test_only_typed_fields_emitted(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(
CompileConfig(enabled=True, backend="inductor", fullgraph=True)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["enable_torch_compile"] is True
assert args.kwargs["torch_compile_kwargs"] == {
"backend": "inductor",
"fullgraph": True,
}
def test_extras_merged_into_torch_compile_kwargs(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(
CompileConfig(
enabled=True,
mode="reduce-overhead",
extras={"options": {"triton.cudagraphs": False}},
)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["torch_compile_kwargs"] == {
"mode": "reduce-overhead",
"options": {"triton.cudagraphs": False},
}
def test_none_fields_suppressed(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["torch_compile_kwargs"] == {}
class TestLegacyLtx2VaeTilingTranslation:
"""``ltx2_vae_tiling`` flat kwarg promotes to
``generator.pipeline.vae_tiling``; reverse direction emits the
legacy name back to FastVideoArgs."""
def test_forward_routes_to_pipeline_vae_tiling(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"ltx2_vae_tiling": False},
)
assert config.pipeline.vae_tiling is False
def test_true_round_trips(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"ltx2_vae_tiling": True},
)
assert config.pipeline.vae_tiling is True
def test_unset_stays_none(self) -> None:
config = legacy_from_pretrained_to_config("/models/ltx2", {})
assert config.pipeline.vae_tiling is None
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
config.pipeline.vae_tiling = False
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["ltx2_vae_tiling"] is False
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
args = generator_config_to_fastvideo_args(config)
assert "ltx2_vae_tiling" not in args.kwargs
class TestLegacyTextEncoderCompileTranslation:
"""``enable_torch_compile_text_encoder`` flat kwarg promotes to
``generator.engine.compile.text_encoder_enabled``; reverse direction
emits the legacy name back onto the FastVideoArgs kwargs dict so
realtime-runtime consumers can read it before FastVideoArgs filters
unknown fields."""
def test_forward_routes_to_compile_text_encoder_enabled(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"enable_torch_compile_text_encoder": True},
)
assert config.engine.compile.text_encoder_enabled is True
def test_false_round_trips(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/ltx2",
{"enable_torch_compile_text_encoder": False},
)
assert config.engine.compile.text_encoder_enabled is False
def test_unset_stays_none(self) -> None:
config = legacy_from_pretrained_to_config("/models/ltx2", {})
assert config.engine.compile.text_encoder_enabled is None
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(
CompileConfig(text_encoder_enabled=True)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["enable_torch_compile_text_encoder"] is True
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/ltx2",
engine=_engine_with_compile(CompileConfig()),
)
args = generator_config_to_fastvideo_args(config)
assert "enable_torch_compile_text_encoder" not in args.kwargs
# -------------------------------------------------------------------
# Helpers
# -------------------------------------------------------------------
def _engine_with_compile(compile_config):
"""Build an ``EngineConfig`` that carries the supplied compile block."""
from fastvideo.api.schema import EngineConfig
engine = EngineConfig()
engine.compile = compile_config
return engine
def _stub_fastvideo_args_from_kwargs(monkeypatch):
"""Swap ``FastVideoArgs.from_kwargs`` for a capture-only stub so
translation tests don't need to construct a valid FastVideoArgs."""
from fastvideo import fastvideo_args as fva
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _Captured)
@@ -0,0 +1,250 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the typed LTX-2 continuation state.
Covers:
* round-trip through :class:`ContinuationState` (inline and blob-backed)
* payload is JSON-serializable (Dynamo RPC / HTTP client constraint)
* kind / schema_version validation on deserialization
* compat-layer validation (known kinds, payload shape)
* round-trip through :func:`request_to_sampling_param` attaches the
state to the resulting :class:`SamplingParam` without losing fidelity
"""
from __future__ import annotations
import json
import numpy as np
import pytest
import torch
# Importing compat first, then the LTX-2 module, exercises the
# self-registration side effect on import (important for the API
# test suite where the pipeline package isn't otherwise imported).
from fastvideo.api import compat as api_compat # noqa: F401
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
OutputConfig,
)
from fastvideo.entrypoints.streaming.session_store import InMemoryBlobStore
from fastvideo.pipelines.basic.ltx2.continuation import (
LTX2_CONTINUATION_KIND,
LTX2_CONTINUATION_SCHEMA_VERSION,
LTX2ContinuationState,
)
def _make_typed_state() -> LTX2ContinuationState:
return LTX2ContinuationState(
segment_index=3,
video_frames=[
(np.ones((64, 64, 3), dtype=np.uint8) * (i * 10)) for i in range(4)
],
video_conditioning_frame_idx=9,
video_conditioning_strength=0.75,
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.float32),
audio_sample_rate=24000,
audio_conditioning_num_frames=5,
audio_conditioning_strength=0.5,
video_position_offset_sec=0.125,
metadata={"note": "unit-test"},
)
class TestRoundTrip:
"""Round-trip through :class:`ContinuationState` preserves all fields."""
def test_kind_and_schema_version(self):
state = _make_typed_state().to_continuation_state()
assert state.kind == LTX2_CONTINUATION_KIND
assert state.payload["schema_version"] == LTX2_CONTINUATION_SCHEMA_VERSION
def test_inline_roundtrip_preserves_scalars(self):
original = _make_typed_state()
envelope = original.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.segment_index == original.segment_index
assert restored.video_conditioning_frame_idx == (
original.video_conditioning_frame_idx)
assert restored.video_conditioning_strength == (
original.video_conditioning_strength)
assert restored.audio_sample_rate == original.audio_sample_rate
assert restored.audio_conditioning_num_frames == (
original.audio_conditioning_num_frames)
assert restored.audio_conditioning_strength == (
original.audio_conditioning_strength)
assert restored.video_position_offset_sec == (
original.video_position_offset_sec)
assert restored.metadata == original.metadata
def test_inline_roundtrip_preserves_video_frames(self):
original = _make_typed_state()
envelope = original.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.video_frames is not None
assert len(restored.video_frames) == len(original.video_frames)
for before, after in zip(original.video_frames,
restored.video_frames):
np.testing.assert_array_equal(before, after)
def test_inline_roundtrip_preserves_audio_latents(self):
original = _make_typed_state()
envelope = original.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.audio_latents is not None
assert tuple(restored.audio_latents.shape) == tuple(
original.audio_latents.shape)
assert restored.audio_latents.dtype == original.audio_latents.dtype
torch.testing.assert_close(
restored.audio_latents, original.audio_latents)
def test_payload_is_json_serializable(self):
envelope = _make_typed_state().to_continuation_state()
# json.dumps must not raise — required for Dynamo RPC transport
# and HTTP client round-trip.
reserialized = json.loads(json.dumps(envelope.payload))
restored = LTX2ContinuationState.from_continuation_state(
ContinuationState(
kind=envelope.kind,
payload=reserialized,
))
assert restored.segment_index == 3
def test_bf16_audio_latents_preserved(self):
"""safetensors serialization must preserve bf16 dtype (numpy
has no bf16, so a raw-bytes path would silently promote)."""
state = LTX2ContinuationState(
segment_index=0,
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.bfloat16),
)
envelope = state.to_continuation_state()
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.audio_latents is not None
assert restored.audio_latents.dtype == torch.bfloat16
torch.testing.assert_close(
restored.audio_latents, state.audio_latents)
class TestBlobIndirection:
"""Large tensors live in the :class:`BlobStore` instead of the payload."""
def test_threshold_triggers_blob_path(self):
blob_store = InMemoryBlobStore()
state = _make_typed_state()
envelope = state.to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=0,
)
assert "blob_id" in envelope.payload["video"]
assert "blob_id" in envelope.payload["audio"]
assert "frames_b64" not in envelope.payload["video"]
assert "safetensors_b64" not in envelope.payload["audio"]
assert len(blob_store) == 2
def test_blob_roundtrip_reconstructs_tensors(self):
blob_store = InMemoryBlobStore()
original = _make_typed_state()
envelope = original.to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=0,
)
restored = LTX2ContinuationState.from_continuation_state(
envelope, blob_store=blob_store)
assert restored.video_frames is not None
assert len(restored.video_frames) == len(original.video_frames)
torch.testing.assert_close(
restored.audio_latents, original.audio_latents)
def test_blob_id_held_when_store_unavailable(self):
"""Deserializing without a blob store preserves the blob id so
the caller can fetch it later."""
blob_store = InMemoryBlobStore()
envelope = _make_typed_state().to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=0,
)
blob_id_video = envelope.payload["video"]["blob_id"]
blob_id_audio = envelope.payload["audio"]["blob_id"]
restored = LTX2ContinuationState.from_continuation_state(envelope)
assert restored.video_frames is None
assert restored.video_frames_blob_id == blob_id_video
assert restored.audio_latents is None
assert restored.audio_latents_blob_id == blob_id_audio
def test_large_threshold_keeps_payload_inline(self):
blob_store = InMemoryBlobStore()
envelope = _make_typed_state().to_continuation_state(
blob_store=blob_store,
inline_threshold_bytes=10 * 1024 * 1024, # 10 MiB
)
assert "frames_b64" in envelope.payload["video"]
assert "safetensors_b64" in envelope.payload["audio"]
assert len(blob_store) == 0
class TestValidation:
"""Invalid payloads error cleanly."""
def test_wrong_kind_rejected(self):
envelope = ContinuationState(kind="longcat.v1", payload={})
with pytest.raises(ValueError, match="Expected ContinuationState.kind"):
LTX2ContinuationState.from_continuation_state(envelope)
def test_unsupported_schema_version_rejected(self):
envelope = ContinuationState(
kind=LTX2_CONTINUATION_KIND,
payload={"schema_version": 999},
)
with pytest.raises(ValueError,
match="Unsupported LTX-2 continuation schema"):
LTX2ContinuationState.from_continuation_state(envelope)
def test_non_png_frame_rejected(self):
state = LTX2ContinuationState(
video_frames=[np.ones((64, 64, 3), dtype=np.float32)],
)
with pytest.raises(ValueError, match="uint8 HxWx3"):
state.to_continuation_state()
class TestCompatLayerWireUp:
"""The public compat layer accepts request.state without reverting
to NotImplementedError and attaches it to the SamplingParam path."""
def test_request_with_state_passes_through(self, tmp_path):
# PR 7 removes the NotImplementedError for request.state; build a
# minimal GenerationRequest carrying an LTX-2 state and make sure
# the public boundary accepts it.
from fastvideo.api.compat import (
normalize_generation_request,
_validate_continuation_state,
)
envelope = _make_typed_state().to_continuation_state()
request = GenerationRequest(
prompt="test",
state=envelope,
)
normalized = normalize_generation_request(request)
_validate_continuation_state(normalized.state)
def test_unknown_kind_rejected_at_boundary(self):
from fastvideo.api.compat import _validate_continuation_state
with pytest.raises(ValueError, match="Unknown ContinuationState kind"):
_validate_continuation_state(
ContinuationState(kind="mystery.v1", payload={}))
def test_empty_kind_rejected_at_boundary(self):
from fastvideo.api.compat import _validate_continuation_state
with pytest.raises(ValueError, match="non-empty string"):
_validate_continuation_state(
ContinuationState(kind="", payload={}))
def test_output_return_state_flag(self):
request = GenerationRequest(
prompt="x",
output=OutputConfig(return_state=True),
)
# The typed public surface exposes the flag directly.
assert request.output.return_state is True
@@ -0,0 +1,298 @@
# SPDX-License-Identifier: Apache-2.0
"""gpu_pool-style flat-kwarg integration tests.
Mirrors the ``load_kwargs`` dict that the FastVideo-internal
``ui/ltx2-streaming/server/gpu_pool.py`` passes to
``VideoGenerator.from_pretrained(**load_kwargs)`` and asserts that the
public typed ``GeneratorConfig`` surface (introduced across PRs 0-6)
can represent it end-to-end, with no fields silently falling through
to ``pipeline.experimental``.
This is the parity guard PR 7.6 depends on: the public gpu_pool
upstream must be able to construct a typed ``GeneratorConfig`` without
knowing any legacy LTX-2 kwarg name, and downstream Dynamo
(``FastVideoArgGroup``) must be able to do the same.
"""
from __future__ import annotations
from copy import deepcopy
import pytest
from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
legacy_from_pretrained_to_config,
)
# Mirrors FastVideo-internal/ui/ltx2-streaming/server/gpu_pool.py
# :lines 233-260 (load_kwargs constructed for VideoGenerator.from_pretrained).
#
# One item from gpu_pool.py's load_kwargs is deliberately excluded:
# - ``pipeline_config=<PipelineConfig instance>`` — an opaque Python
# object; internal mutates it in place (``dit_config.quant_config =
# FP4Config()``). The typed path for quantization is tracked in
# "Known Technical Debt" in PR plan.md; ``pipeline_config`` as an
# instance legitimately belongs in ``pipeline.experimental``.
#
# ``enable_torch_compile_text_encoder`` IS included below: its typed
# home is ``CompileConfig.text_encoder_enabled`` (added post-review).
# The legacy ``FastVideoArgs`` path does not yet consume it; the
# realtime runtime (PR 7.6) reads it off the kwargs dict before
# FastVideoArgs filtering.
GPU_POOL_LOAD_KWARGS = {
"config_model_path": "/models/ltx2-distilled/config",
"num_gpus": 1,
"dit_layerwise_offload": False,
"use_fsdp_inference": False,
"dit_cpu_offload": False,
"vae_cpu_offload": False,
"text_encoder_cpu_offload": False,
"pin_cpu_memory": True,
"ltx2_vae_tiling": False,
"ltx2_refine_enabled": True,
"ltx2_refine_upsampler_path": "/models/ltx2-distilled/spatial_upsampler",
"ltx2_refine_lora_path": "",
"ltx2_refine_num_inference_steps": 2,
"ltx2_refine_guidance_scale": 1.0,
"ltx2_refine_add_noise": True,
"enable_torch_compile": True,
"enable_torch_compile_text_encoder": True,
"torch_compile_kwargs": {
"backend": "inductor",
"fullgraph": True,
"mode": "max-autotune-no-cudagraphs",
"dynamic": False,
},
}
class TestGpuPoolForwardTranslation:
"""gpu_pool flat kwargs -> typed GeneratorConfig."""
@pytest.fixture(scope="class")
def config(self):
return legacy_from_pretrained_to_config(
"FastVideo/LTX2-Distilled-Diffusers",
GPU_POOL_LOAD_KWARGS,
)
def test_model_path_set(self, config) -> None:
assert config.model_path == "FastVideo/LTX2-Distilled-Diffusers"
def test_engine_basics(self, config) -> None:
assert config.engine.num_gpus == 1
assert config.engine.use_fsdp_inference is False
def test_offload_config(self, config) -> None:
assert config.engine.offload.dit is False
assert config.engine.offload.dit_layerwise is False
assert config.engine.offload.vae is False
assert config.engine.offload.text_encoder is False
assert config.engine.offload.pin_cpu_memory is True
def test_compile_config_typed_fields_extracted(self, config) -> None:
compile_config = config.engine.compile
assert compile_config.enabled is True
assert compile_config.text_encoder_enabled is True
assert compile_config.backend == "inductor"
assert compile_config.fullgraph is True
assert compile_config.mode == "max-autotune-no-cudagraphs"
assert compile_config.dynamic is False
assert compile_config.extras == {}
def test_vae_tiling_routed_to_pipeline(self, config) -> None:
assert config.pipeline.vae_tiling is False
def test_config_model_path_routed_to_components(self, config) -> None:
assert config.pipeline.components.config_root == "/models/ltx2-distilled/config"
def test_refine_upsampler_routed_to_components(self, config) -> None:
assert config.pipeline.components.upsampler_weights == (
"/models/ltx2-distilled/spatial_upsampler")
def test_empty_refine_lora_becomes_none(self, config) -> None:
# gpu_pool passes "" to keep refine LoRA disabled; typed schema
# treats that as "no LoRA" rather than an empty-string path.
assert config.pipeline.components.lora_path is None
def test_refine_preset_overrides(self, config) -> None:
refine = config.pipeline.preset_overrides.get("refine", {})
assert refine == {
"enabled": True,
"num_inference_steps": 2,
"guidance_scale": 1.0,
"add_noise": True,
}
def test_no_experimental_leakage(self, config) -> None:
"""Every gpu_pool kwarg should have a typed home — nothing should
silently fall through to ``pipeline.experimental``."""
assert config.pipeline.experimental == {}
class TestGpuPoolReverseTranslation:
"""typed GeneratorConfig -> FastVideoArgs kwargs reproduces the
original gpu_pool flat-kwarg shape.
This is what lets PR 7.6 wire the public ``gpu_pool`` through
``generator_config_to_fastvideo_args`` without the runtime noticing.
"""
@pytest.fixture
def args_kwargs(self, monkeypatch):
from fastvideo import fastvideo_args as fva
captured: dict[str, object] = {}
def _capture(**kw):
captured.update(kw)
return _Captured(**kw)
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
config = legacy_from_pretrained_to_config(
"FastVideo/LTX2-Distilled-Diffusers",
GPU_POOL_LOAD_KWARGS,
)
generator_config_to_fastvideo_args(config)
return captured
def test_ltx2_refine_flags_reemitted(self, args_kwargs) -> None:
assert args_kwargs["ltx2_refine_enabled"] is True
assert args_kwargs["ltx2_refine_add_noise"] is True
assert args_kwargs["ltx2_refine_num_inference_steps"] == 2
assert args_kwargs["ltx2_refine_guidance_scale"] == 1.0
def test_refine_upsampler_path_reemitted(self, args_kwargs) -> None:
assert args_kwargs["ltx2_refine_upsampler_path"] == (
"/models/ltx2-distilled/spatial_upsampler")
def test_config_model_path_reemitted(self, args_kwargs) -> None:
assert args_kwargs["config_model_path"] == "/models/ltx2-distilled/config"
def test_torch_compile_kwargs_reassembled(self, args_kwargs) -> None:
assert args_kwargs["torch_compile_kwargs"] == {
"backend": "inductor",
"fullgraph": True,
"mode": "max-autotune-no-cudagraphs",
"dynamic": False,
}
def test_vae_tiling_reemitted_with_legacy_name(self, args_kwargs) -> None:
assert args_kwargs["ltx2_vae_tiling"] is False
def test_text_encoder_compile_reemitted(self, args_kwargs) -> None:
# Present in the captured kwargs dict even though
# ``FastVideoArgs.from_kwargs`` will filter it out — realtime
# runtime upstream (PR 7.6) reads it off this dict.
assert args_kwargs["enable_torch_compile_text_encoder"] is True
def test_no_stray_refine_dict(self, args_kwargs) -> None:
"""preset_overrides.refine must flatten to ltx2_refine_* kwargs
rather than landing as a nested ``refine`` kwarg that
FastVideoArgs doesn't understand."""
assert "refine" not in args_kwargs
class TestRefineFlattenCoversAllTypedFields:
"""Every field on LTX2Refine{Preset,Stage}Override must survive the
round-trip through preset_overrides.refine back to ltx2_refine_*
kwargs. Guards against the hardcoded-key-tuple regression where
image_crf / video_position_offset_sec silently dropped."""
def test_all_fields_reemitted(self, monkeypatch) -> None:
from fastvideo import fastvideo_args as fva
from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
)
from fastvideo.api.schema import GeneratorConfig, PipelineSelection
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_preset_override_fields,
refine_stage_override_fields,
)
captured: dict[str, object] = {}
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
def _capture(**kw):
captured.update(kw)
return _Captured(**kw)
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
refine_payload = {
# Preset-override fields.
"enabled": True,
"add_noise": False,
# Stage-override fields.
"num_inference_steps": 3,
"guidance_scale": 1.5,
"image_crf": 18,
"video_position_offset_sec": 2.5,
}
all_fields = (refine_preset_override_fields()
| refine_stage_override_fields())
assert set(refine_payload) == all_fields, (
"payload must cover every typed field to exercise the flatten loop")
config = GeneratorConfig(
model_path="/models/ltx2",
pipeline=PipelineSelection(preset_overrides={"refine": refine_payload}),
)
generator_config_to_fastvideo_args(config)
for key, value in refine_payload.items():
assert captured[f"ltx2_refine_{key}"] == value
class TestCompileExtrasPreserved:
"""Additional torch.compile kwargs beyond the four typed fields
round-trip through ``CompileConfig.extras``."""
def test_extras_preserved(self, monkeypatch) -> None:
from fastvideo import fastvideo_args as fva
captured: dict[str, object] = {}
def _capture(**kw):
captured.update(kw)
class _Captured:
def __init__(self, **kw):
self.kwargs = kw
return _Captured(**kw)
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
kwargs = deepcopy(GPU_POOL_LOAD_KWARGS)
kwargs["torch_compile_kwargs"] = {
"backend": "inductor",
"options": {"triton.cudagraphs": False},
"disable": False,
}
config = legacy_from_pretrained_to_config(
"FastVideo/LTX2-Distilled-Diffusers", kwargs)
assert config.engine.compile.backend == "inductor"
assert config.engine.compile.extras == {
"options": {"triton.cudagraphs": False},
"disable": False,
}
generator_config_to_fastvideo_args(config)
assert captured["torch_compile_kwargs"] == {
"backend": "inductor",
"options": {"triton.cudagraphs": False},
"disable": False,
}
@@ -0,0 +1,121 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for typed LTX-2 stage override dataclasses."""
from __future__ import annotations
import pytest
from fastvideo.api.errors import ConfigValidationError
from fastvideo.api.presets import get_preset, validate_stage_overrides
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
LTX2RefinePresetOverride,
LTX2RefineStageOverride,
refine_override_to_dict,
refine_preset_override_fields,
refine_stage_override_fields,
)
class TestRefineStageOverrideDataclass:
def test_all_fields_default_to_none(self) -> None:
override = LTX2RefineStageOverride()
assert override.num_inference_steps is None
assert override.guidance_scale is None
assert override.image_crf is None
assert override.video_position_offset_sec is None
def test_explicit_construction(self) -> None:
override = LTX2RefineStageOverride(
num_inference_steps=2,
guidance_scale=1.0,
image_crf=18,
video_position_offset_sec=2.5,
)
assert override.num_inference_steps == 2
assert override.guidance_scale == 1.0
assert override.image_crf == 18
assert override.video_position_offset_sec == 2.5
def test_to_dict_drops_none(self) -> None:
override = LTX2RefineStageOverride(num_inference_steps=3)
assert refine_override_to_dict(override) == {
"num_inference_steps": 3,
}
def test_to_dict_with_all_fields(self) -> None:
override = LTX2RefineStageOverride(
num_inference_steps=2,
guidance_scale=1.0,
image_crf=18,
video_position_offset_sec=0.0,
)
assert refine_override_to_dict(override) == {
"num_inference_steps": 2,
"guidance_scale": 1.0,
"image_crf": 18,
"video_position_offset_sec": 0.0,
}
def test_fields_accessor_matches_dataclass(self) -> None:
assert refine_stage_override_fields() == frozenset({
"num_inference_steps",
"guidance_scale",
"image_crf",
"video_position_offset_sec",
})
class TestRefinePresetOverrideDataclass:
def test_all_fields_default_to_none(self) -> None:
override = LTX2RefinePresetOverride()
assert override.enabled is None
assert override.add_noise is None
def test_to_dict_drops_none(self) -> None:
override = LTX2RefinePresetOverride(enabled=True)
assert refine_override_to_dict(override) == {
"enabled": True,
}
def test_to_dict_with_all_fields(self) -> None:
override = LTX2RefinePresetOverride(enabled=True, add_noise=False)
assert refine_override_to_dict(override) == {
"enabled": True,
"add_noise": False,
}
def test_fields_accessor_matches_dataclass(self) -> None:
assert refine_preset_override_fields() == frozenset({
"enabled",
"add_noise",
})
class TestStageOverridesMirrorPresetSchema:
"""The ltx2_two_stage preset's refine stage schema must list
exactly the :class:`LTX2RefineStageOverride` field names."""
def test_allowed_overrides_mirror_dataclass(self) -> None:
import fastvideo.registry # noqa: F401
preset = get_preset("ltx2_two_stage", "ltx2")
refine_schema = next(
s for s in preset.stage_schemas if s.name == "refine")
assert refine_schema.allowed_overrides == refine_stage_override_fields()
def test_roundtrip_through_validate_stage_overrides(self) -> None:
import fastvideo.registry # noqa: F401
preset = get_preset("ltx2_two_stage", "ltx2")
override = LTX2RefineStageOverride(
num_inference_steps=3,
guidance_scale=1.0,
)
validate_stage_overrides(
preset, {"refine": refine_override_to_dict(override)})
def test_unknown_field_rejected(self) -> None:
import fastvideo.registry # noqa: F401
preset = get_preset("ltx2_two_stage", "ltx2")
with pytest.raises(ConfigValidationError):
validate_stage_overrides(
preset, {"refine": {"unknown_key": 1}})
+10 -1
View File
@@ -111,7 +111,15 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
"vae": True,
"pin_cpu_memory": True,
},
"compile": {"enabled": False, "kwargs": {}},
"compile": {
"enabled": False,
"text_encoder_enabled": None,
"backend": None,
"fullgraph": None,
"mode": None,
"dynamic": None,
"extras": {},
},
"enable_stage_verification": True,
"use_fsdp_inference": False,
"disable_autocast": False,
@@ -133,6 +141,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
"override_pipeline_cls_name": None,
"override_transformer_cls_name": None,
},
"vae_tiling": None,
"preset_overrides": {},
"experimental": {},
},
+75 -1
View File
@@ -334,7 +334,7 @@ class TestLtx2Presets:
import fastvideo.registry # noqa: F401
presets = get_presets_for_family("ltx2")
names = {p.name for p in presets}
assert names == {"ltx2_base", "ltx2_distilled"}
assert names == {"ltx2_base", "ltx2_distilled", "ltx2_two_stage"}
def test_ltx2_base_lookup(self) -> None:
import fastvideo.registry # noqa: F401
@@ -349,6 +349,43 @@ class TestLtx2Presets:
assert p.defaults["num_inference_steps"] == 8
assert p.defaults["guidance_scale"] == 1.0
def test_ltx2_two_stage_is_two_stage(self) -> None:
import fastvideo.registry # noqa: F401
p = get_preset("ltx2_two_stage", "ltx2")
assert len(p.stage_schemas) == 2
assert p.stage_schemas[0].name == "denoise"
assert p.stage_schemas[1].name == "refine"
assert p.stage_schemas[1].kind == "refinement"
def test_ltx2_two_stage_stage_defaults(self) -> None:
import fastvideo.registry # noqa: F401
p = get_preset("ltx2_two_stage", "ltx2")
refine = p.stage_defaults["refine"]
# stage-2 refine only supports 2 or 3 denoising steps; preset
# defaults to 2 (matches gpu_pool.py load_kwargs).
assert refine["num_inference_steps"] == 2
assert refine["guidance_scale"] == 1.0
def test_ltx2_two_stage_refine_overrides_valid(self) -> None:
import fastvideo.registry # noqa: F401
p = get_preset("ltx2_two_stage", "ltx2")
validate_stage_overrides(
p, {"refine": {"num_inference_steps": 3}})
validate_stage_overrides(
p, {"refine": {"guidance_scale": 1.0}})
validate_stage_overrides(
p, {"refine": {"image_crf": 18}})
validate_stage_overrides(
p, {"refine": {"video_position_offset_sec": 2.5}})
def test_ltx2_two_stage_rejects_unknown_refine_override(self) -> None:
import fastvideo.registry # noqa: F401
from fastvideo.api.errors import ConfigValidationError
p = get_preset("ltx2_two_stage", "ltx2")
with pytest.raises(ConfigValidationError):
validate_stage_overrides(
p, {"refine": {"bogus_field": 1}})
# -------------------------------------------------------------------
# Hunyuan preset integration
@@ -548,3 +585,40 @@ class TestPresetCountIntegrity:
import fastvideo.registry # noqa: F401
names = get_all_preset_names()
assert len(names) >= 37
class TestPresetDefaultTypes:
"""Preset ``defaults`` values must match the types on
:class:`SamplingParam`. Assigning ``None`` to a typed-``str`` field
(e.g. ``negative_prompt``) breaks downstream stages that assert the
runtime type — see the CFG branch in
``pipelines/stages/text_encoding.py:81``."""
def test_ltx2_cfg_defaults_are_off(self) -> None:
"""SamplingParam's LTX-2 CFG class defaults must be 1.0 (CFG
off). ``ForwardBatch.__post_init__`` force-enables
``do_classifier_free_guidance`` when either
``ltx2_cfg_scale_video`` or ``ltx2_cfg_scale_audio`` is != 1.0,
so any non-1.0 default silently forces CFG on for every model
family that doesn't explicitly override these fields. Guard
against the regression that surfaced as the TurboDiffusion I2V
SSIM crash (``text_encoding.py:81`` assertion on
``negative_prompt``)."""
from fastvideo.api.sampling_param import SamplingParam
sp = SamplingParam()
assert sp.ltx2_cfg_scale_video == 1.0
assert sp.ltx2_cfg_scale_audio == 1.0
def test_no_preset_sets_negative_prompt_to_none(self) -> None:
import fastvideo.registry # noqa: F401
from fastvideo.api.presets import _PRESET_REGISTRY
offenders = [
f"{preset.model_family}/{preset.name}"
for preset in _PRESET_REGISTRY.values()
if preset.defaults.get("negative_prompt", "") is None
]
assert not offenders, (
"These presets set negative_prompt=None, which violates "
"SamplingParam.negative_prompt's typed str contract and "
"crashes the CFG path in text_encoding. Use \"\" instead:\n"
+ "\n".join(f" - {p}" for p in offenders))
@@ -46,24 +46,42 @@ def _flatten_status_section(section: dict, valid_statuses: set[str]) -> set[str]
return names
def _get_extra_dataclass_fields(package_name: str, base_cls: type) -> set[str]:
package = importlib.import_module(package_name)
def _get_extra_dataclass_fields(
package_names: str | tuple[str, ...],
base_cls: type,
) -> set[str]:
"""Collect dataclass fields declared on ``base_cls`` subclasses found
under any of the given package roots.
Accepts either a single package name (string) or a tuple of package
roots — the latter supports the PR 6 colocation where each model
family's ``PipelineConfig`` subclass moves from
``fastvideo.configs.pipelines.<family>`` to
``fastvideo.pipelines.basic.<family>.pipeline_configs``.
"""
if isinstance(package_names, str):
package_names = (package_names, )
base_fields = {f.name for f in dataclasses.fields(base_cls)}
extras: set[str] = set()
if not hasattr(package, "__path__"):
return extras
for _, modname, _ in pkgutil.iter_modules(package.__path__):
if modname == "__pycache__":
for package_name in package_names:
package = importlib.import_module(package_name)
if not hasattr(package, "__path__"):
continue
module = importlib.import_module(f"{package_name}.{modname}")
for obj in vars(module).values():
if (
isinstance(obj, type)
and dataclasses.is_dataclass(obj)
and issubclass(obj, base_cls)
and obj is not base_cls
):
extras.update(f.name for f in dataclasses.fields(obj) if f.name not in base_fields)
for _, modname, is_pkg in pkgutil.walk_packages(
package.__path__, prefix=f"{package_name}."):
if modname.endswith(".__pycache__"):
continue
module = importlib.import_module(modname)
for obj in vars(module).values():
if (isinstance(obj, type)
and dataclasses.is_dataclass(obj)
and issubclass(obj, base_cls)
and obj is not base_cls):
extras.update(
f.name for f in dataclasses.fields(obj)
if f.name not in base_fields)
return extras
@@ -177,7 +195,10 @@ def test_pipeline_config_base_fields_are_classified() -> None:
def test_pipeline_config_extension_fields_are_classified() -> None:
inventory = _load_inventory()
expected = _get_extra_dataclass_fields("fastvideo.configs.pipelines", PipelineConfig)
expected = _get_extra_dataclass_fields(
("fastvideo.configs.pipelines", "fastvideo.pipelines.basic"),
PipelineConfig,
)
actual = _flatten_status_section(
inventory["surfaces"]["pipeline_config_extensions"],
set(inventory["status_definitions"]),
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
"""Protocol schema tests for the streaming server.
Covers:
* accepted client messages parse into the correct discriminated model
* unknown ``type`` values raise validation errors
* server-side messages serialize to the expected wire shape
* continuation_state field on session_init_v2 carries through
"""
from __future__ import annotations
import pytest
from pydantic import ValidationError
from fastvideo.entrypoints.streaming.protocol import (
ContinuationStateSnapshot,
ErrorMessage,
GpuAssigned,
Ltx2SegmentComplete,
Ltx2SegmentStart,
Ltx2StreamStart,
MediaInit,
MediaSegmentComplete,
QueueStatus,
SegmentPromptSource,
SessionInitV2,
SnapshotState,
StepComplete,
parse_client_message,
)
class TestClientMessageParsing:
def test_session_init_v2_minimal(self):
parsed = parse_client_message({"type": "session_init_v2"})
assert isinstance(parsed, SessionInitV2)
assert parsed.curated_prompts == []
assert parsed.stream_mode == "av_fmp4"
def test_session_init_v2_full(self):
raw = {
"type": "session_init_v2",
"client_id": "client-1",
"preset": "ltx2_two_stage",
"preset_label": "2x refine",
"curated_prompts": ["a fox", "a deer"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
"single_clip_mode": True,
"stream_mode": "av_fmp4",
"continuation_state": {
"kind": "ltx2.v1",
"payload": {"schema_version": 1, "segment_index": 2},
},
}
parsed = parse_client_message(raw)
assert isinstance(parsed, SessionInitV2)
assert parsed.preset == "ltx2_two_stage"
assert parsed.curated_prompts == ["a fox", "a deer"]
assert parsed.continuation_state["kind"] == "ltx2.v1"
def test_segment_prompt_source(self):
parsed = parse_client_message({
"type": "segment_prompt_source",
"prompt": "hello world",
"source": "curated",
"seed": 7,
})
assert isinstance(parsed, SegmentPromptSource)
assert parsed.source == "curated"
assert parsed.seed == 7
def test_snapshot_state(self):
parsed = parse_client_message({"type": "snapshot_state"})
assert isinstance(parsed, SnapshotState)
def test_unknown_type_rejected(self):
with pytest.raises(ValidationError):
parse_client_message({"type": "not_a_real_message"})
def test_missing_type_rejected(self):
with pytest.raises(ValidationError):
parse_client_message({"prompt": "x"})
def test_segment_prompt_source_requires_prompt(self):
with pytest.raises(ValidationError):
parse_client_message({"type": "segment_prompt_source"})
class TestServerMessageSerialization:
def test_queue_status(self):
msg = QueueStatus(position=3, queue_depth=5)
assert msg.model_dump() == {
"type": "queue_status",
"position": 3,
"queue_depth": 5,
}
def test_gpu_assigned(self):
msg = GpuAssigned(gpu_id=1, session_timeout=300)
assert msg.model_dump()["type"] == "gpu_assigned"
def test_ltx2_stream_start(self):
msg = Ltx2StreamStart(
preset="ltx2_two_stage",
width=1024, height=1536, fps=24, num_frames=121,
)
dumped = msg.model_dump()
assert dumped["type"] == "ltx2_stream_start"
assert dumped["width"] == 1024
def test_ltx2_segment_start(self):
msg = Ltx2SegmentStart(
segment_idx=0,
prompt="a fox",
total_steps=8,
)
assert msg.model_dump()["segment_idx"] == 0
def test_step_complete(self):
msg = StepComplete(segment_idx=0, step=1, total_steps=8)
assert msg.model_dump()["stage"] == "denoise"
def test_media_init_has_mode(self):
msg = MediaInit(segment_idx=0, stream_id="abc")
dumped = msg.model_dump()
assert dumped["mode"] == "av_fmp4"
assert "avc1" in dumped["mime"]
def test_media_segment_complete(self):
msg = MediaSegmentComplete(
segment_idx=0, stream_id="abc", chunks=4,
)
dumped = msg.model_dump()
assert dumped["chunks"] == 4
def test_ltx2_segment_complete(self):
msg = Ltx2SegmentComplete(segment_idx=0, generation_time_ms=1234.5)
assert msg.model_dump()["generation_time_ms"] == 1234.5
def test_error_message_code_restricted(self):
with pytest.raises(ValidationError):
ErrorMessage(code="not_a_code", message="x")
def test_continuation_state_snapshot(self):
msg = ContinuationStateSnapshot(state={
"kind": "ltx2.v1",
"payload": {"schema_version": 1},
})
assert msg.model_dump()["state"]["kind"] == "ltx2.v1"
@@ -0,0 +1,237 @@
# SPDX-License-Identifier: Apache-2.0
"""End-to-end WebSocket smoke for the streaming server skeleton.
Uses a mock generator so these tests run CPU-only (no GPU, no model
weights). Skips the fMP4 assertions when ``ffmpeg`` is missing.
"""
from __future__ import annotations
import shutil
from dataclasses import dataclass
from typing import Any
import numpy as np
import pytest
pytest.importorskip("starlette")
from starlette.testclient import TestClient # noqa: E402
from fastvideo.api.schema import ( # noqa: E402
ContinuationState,
GeneratorConfig,
SamplingConfig,
ServeConfig,
StreamingConfig,
GenerationRequest,
)
from fastvideo.entrypoints.streaming.server import build_app # noqa: E402
_FFMPEG_AVAILABLE = shutil.which("ffmpeg") is not None
@dataclass
class _MockGenerator:
width: int = 64
height: int = 64
fps: int = 12
num_frames: int = 12
return_state: bool = True
def generate(self, request: GenerationRequest) -> dict[str, Any]:
frames = [
np.full((self.height, self.width, 3), i * 5, dtype=np.uint8)
for i in range(self.num_frames)
]
state = (ContinuationState(
kind="ltx2.v1",
payload={
"schema_version": 1,
"segment_index": 0,
"source_prompt": request.prompt,
},
) if self.return_state else None)
return {
"frames": frames,
"audio_sample_rate": 24000,
"state": state,
}
def _build_serve_config() -> ServeConfig:
return ServeConfig(
generator=GeneratorConfig(model_path="/models/fake"),
default_request=GenerationRequest(
sampling=SamplingConfig(
num_frames=12,
height=64,
width=64,
fps=12,
num_inference_steps=1,
),
),
streaming=StreamingConfig(
session_timeout_seconds=60,
generation_segment_cap=2,
),
)
def _build_client() -> tuple[TestClient, _MockGenerator]:
generator = _MockGenerator()
app = build_app(_build_serve_config(), generator)
return TestClient(app), generator
class TestHealth:
def test_health_endpoint_reports_stream_mode(self):
client, _ = _build_client()
response = client.get("/health")
assert response.status_code == 200
body = response.json()
assert body["status"] == "ok"
assert body["stream_mode"] == "av_fmp4"
assert body["sessions"] == 0
class TestSessionHandshake:
def test_rejects_non_session_init_opening_frame(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "segment_prompt_source", "prompt": "x"})
err = ws.receive_json()
assert err["type"] == "error"
assert err["code"] == "invalid_message"
def test_rejects_unknown_message_on_init(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "not_a_message"})
err = ws.receive_json()
assert err["type"] == "error"
def test_emits_queue_and_gpu_assigned_on_valid_init(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({
"type": "session_init_v2",
"preset": "ltx2_two_stage",
"curated_prompts": ["a fox"],
})
assert ws.receive_json()["type"] == "queue_status"
assert ws.receive_json()["type"] == "gpu_assigned"
assert ws.receive_json()["type"] == "ltx2_stream_start"
def test_init_hydrates_continuation_state(self):
client, _ = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({
"type": "session_init_v2",
"preset": "ltx2_two_stage",
"continuation_state": {
"kind": "ltx2.v1",
"payload": {"schema_version": 1, "segment_index": 3},
},
})
# Drain handshake frames
ws.receive_json() # queue_status
ws.receive_json() # gpu_assigned
ws.receive_json() # ltx2_stream_start
# Ask the server for the state back; it should echo what we sent.
ws.send_json({"type": "snapshot_state"})
snap = ws.receive_json()
assert snap["type"] == "continuation_state_snapshot"
assert snap["state"]["kind"] == "ltx2.v1"
assert snap["state"]["payload"]["segment_index"] == 3
@pytest.mark.skipif(not _FFMPEG_AVAILABLE, reason="ffmpeg not installed")
class TestSegmentFlow:
def test_segment_generates_media_init_plus_complete(self):
client, generator = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "session_init_v2",
"preset": "ltx2_two_stage"})
for _ in range(3):
ws.receive_json() # queue_status + gpu_assigned + stream_start
ws.send_json({
"type": "segment_prompt_source",
"prompt": "a test segment",
"num_inference_steps": 1,
})
start = ws.receive_json()
assert start["type"] == "ltx2_segment_start"
assert start["segment_idx"] == 0
step = ws.receive_json()
assert step["type"] == "step_complete"
media_init = ws.receive_json()
assert media_init["type"] == "media_init"
# Then one or more binary frames until media_segment_complete.
saw_binary = False
while True:
msg = ws.receive()
if "bytes" in msg and msg["bytes"]:
saw_binary = True
continue
parsed = _as_json(msg)
if parsed is None:
continue
if parsed["type"] == "media_segment_complete":
break
assert saw_binary
final = ws.receive_json()
assert final["type"] == "ltx2_segment_complete"
assert final["segment_idx"] == 0
class TestContinuationStatePersistence:
def test_snapshot_after_segment_carries_generator_state(self):
if not _FFMPEG_AVAILABLE:
pytest.skip("ffmpeg not installed")
client, generator = _build_client()
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "session_init_v2",
"preset": "ltx2_two_stage"})
for _ in range(3):
ws.receive_json()
ws.send_json({
"type": "segment_prompt_source",
"prompt": "a cat",
"num_inference_steps": 1,
})
_drain_until(ws, "ltx2_segment_complete")
ws.send_json({"type": "snapshot_state"})
snap = ws.receive_json()
assert snap["type"] == "continuation_state_snapshot"
assert snap["state"]["kind"] == "ltx2.v1"
assert snap["state"]["payload"]["source_prompt"] == "a cat"
# ----------------------------------------------------------------------
# Helpers
# ----------------------------------------------------------------------
def _drain_until(ws, target_type: str) -> dict[str, Any]:
while True:
msg = ws.receive()
if "text" in msg and msg["text"]:
import json
parsed = json.loads(msg["text"])
if parsed.get("type") == target_type:
return parsed
# skip binary / other
def _as_json(msg: dict[str, Any]) -> dict[str, Any] | None:
if "text" not in msg or not msg["text"]:
return None
import json
return json.loads(msg["text"])
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
"""Session lifecycle tests."""
from __future__ import annotations
import time
import pytest
from fastvideo.entrypoints.streaming.session import (
InvalidSessionTransition,
Session,
SessionManager,
SessionRejected,
SessionState,
)
class TestSessionStateMachine:
def test_starts_initializing(self):
s = Session()
assert s.state is SessionState.INITIALIZING
def test_legal_sequence(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.transition(SessionState.COMPLETE)
assert s.state is SessionState.COMPLETE
def test_active_self_loop_allowed(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.transition(SessionState.ACTIVE) # re-asserting is fine
assert s.state is SessionState.ACTIVE
def test_illegal_backwards_transition(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
with pytest.raises(InvalidSessionTransition):
s.transition(SessionState.INITIALIZING)
def test_cannot_leave_terminal_state(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.transition(SessionState.COMPLETE)
with pytest.raises(InvalidSessionTransition):
s.transition(SessionState.ACTIVE)
def test_error_terminal(self):
s = Session()
s.transition(SessionState.QUEUED)
s.transition(SessionState.ERROR)
with pytest.raises(InvalidSessionTransition):
s.transition(SessionState.ACTIVE)
def test_transition_updates_activity(self):
s = Session()
prior = s.last_activity
time.sleep(0.001)
s.transition(SessionState.QUEUED)
assert s.last_activity > prior
def test_segment_cap(self):
s = Session()
s.segment_idx = 5
assert s.segment_cap_reached(5) is True
assert s.segment_cap_reached(6) is False
class TestSessionManager:
def test_create_assigns_unique_ids(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=2)
a = mgr.create()
b = mgr.create()
assert a.id != b.id
assert len(mgr) == 2
def test_max_sessions_enforced(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
mgr.create()
with pytest.raises(SessionRejected):
mgr.create()
def test_close_releases_slot(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
s = mgr.create()
mgr.close(s.id)
assert len(mgr) == 0
# Now can create again.
mgr.create()
def test_reap_timed_out_flags_idle_sessions(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
s = mgr.create()
s.transition(SessionState.QUEUED)
s.transition(SessionState.GPU_BINDING)
s.transition(SessionState.ACTIVE)
s.last_activity = time.monotonic() - 10 # 10s ago, past the 1s budget
dead = mgr.reap_timed_out()
assert s.id in dead
def test_reap_skips_terminal_states(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
s = mgr.create()
s.transition(SessionState.QUEUED)
s.transition(SessionState.ERROR)
s.last_activity = time.monotonic() - 10
assert s.id not in mgr.reap_timed_out()
def test_active_sessions_filter(self):
mgr = SessionManager(
segment_cap=4, session_timeout_seconds=60, max_sessions=4)
a = mgr.create()
a.transition(SessionState.QUEUED)
a.transition(SessionState.GPU_BINDING)
a.transition(SessionState.ACTIVE)
b = mgr.create() # INITIALIZING
assert mgr.active_sessions() == [a]
assert b not in mgr.active_sessions()
@@ -0,0 +1,90 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the session init-image persistence helper."""
from __future__ import annotations
import base64
import io
import os
import pytest
from PIL import Image
from fastvideo.entrypoints.streaming.session_init_image import (
persist_session_init_image,
)
def _png_bytes(size: tuple[int, int] = (64, 64)) -> bytes:
buffer = io.BytesIO()
Image.new("RGB", size, color=(10, 20, 30)).save(buffer, format="PNG")
return buffer.getvalue()
class TestPersistSessionInitImage:
def test_none_payload_returns_none(self):
assert persist_session_init_image(None) is None
assert persist_session_init_image({}) is None
def test_non_object_payload_rejected(self):
with pytest.raises(ValueError):
persist_session_init_image("not-a-dict")
def test_png_payload_persists(self, tmp_path):
data = _png_bytes()
image = persist_session_init_image({
"mime": "image/png",
"name": "ref.png",
"data": base64.b64encode(data).decode("ascii"),
}, output_dir=str(tmp_path))
assert image is not None
assert os.path.exists(image.path)
assert image.mime == "image/png"
assert image.path.endswith(".png")
with open(image.path, "rb") as f:
assert f.read() == data
def test_unknown_mime_rejected(self, tmp_path):
with pytest.raises(ValueError, match="mime"):
persist_session_init_image({
"mime": "image/bmp",
"data": "ignored",
}, output_dir=str(tmp_path))
def test_bad_base64_rejected(self, tmp_path):
with pytest.raises(ValueError, match="base64"):
persist_session_init_image({
"mime": "image/png",
"data": "not!base64!",
}, output_dir=str(tmp_path))
def test_empty_data_rejected(self, tmp_path):
with pytest.raises(ValueError, match="empty"):
persist_session_init_image({
"mime": "image/png",
"data": "",
}, output_dir=str(tmp_path))
def test_display_name_sanitized(self, tmp_path):
image = persist_session_init_image({
"mime": "image/png",
"name": "../evil/../name.png",
"data": base64.b64encode(_png_bytes()).decode("ascii"),
}, output_dir=str(tmp_path))
assert image is not None
assert image.display_name == "name.png"
def test_oversize_rejected(self, tmp_path):
from fastvideo.entrypoints.streaming import session_init_image as mod
original = mod._MAX_IMAGE_BYTES
mod._MAX_IMAGE_BYTES = 100
try:
with pytest.raises(ValueError, match="limit"):
persist_session_init_image({
"mime": "image/png",
"data": base64.b64encode(_png_bytes((512, 512))).decode(
"ascii"),
}, output_dir=str(tmp_path))
finally:
mod._MAX_IMAGE_BYTES = original
@@ -0,0 +1,186 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the streaming SessionStore and BlobStore.
Covers:
* ``store`` / ``snapshot`` / ``drop`` lifecycle for the in-memory store
* ``hydrate`` with and without an explicit session id
* blob store insert / get / drop semantics
* thread-safety under concurrent writes (smoke)
* round-trip a LTX-2 continuation through snapshot + hydrate across a
session boundary (the "export and resume" flow the PR plan calls out)
"""
from __future__ import annotations
import threading
import numpy as np
import pytest
import torch
from fastvideo.api.schema import ContinuationState
from fastvideo.entrypoints.streaming.session_store import (
BlobStore,
InMemoryBlobStore,
InMemorySessionStore,
SessionStore,
)
from fastvideo.pipelines.basic.ltx2.continuation import (
LTX2_CONTINUATION_KIND,
LTX2ContinuationState,
)
class TestInMemoryBlobStore:
def test_is_blob_store(self):
assert isinstance(InMemoryBlobStore(), BlobStore)
def test_put_then_get_returns_same_bytes(self):
store = InMemoryBlobStore()
blob_id = store.put(b"hello")
assert store.get(blob_id) == b"hello"
def test_put_returns_distinct_ids(self):
store = InMemoryBlobStore()
id_a = store.put(b"a")
id_b = store.put(b"b")
assert id_a != id_b
def test_get_missing_raises_keyerror(self):
store = InMemoryBlobStore()
with pytest.raises(KeyError):
store.get("nonexistent")
def test_drop_removes_blob(self):
store = InMemoryBlobStore()
blob_id = store.put(b"payload")
store.drop(blob_id)
assert blob_id not in store
with pytest.raises(KeyError):
store.get(blob_id)
def test_drop_missing_is_noop(self):
store = InMemoryBlobStore()
store.drop("not-there") # no raise
def test_contains(self):
store = InMemoryBlobStore()
blob_id = store.put(b"x")
assert blob_id in store
assert "other" not in store
class TestInMemorySessionStore:
def test_is_session_store(self):
assert isinstance(InMemorySessionStore(), SessionStore)
def test_store_then_snapshot(self):
store = InMemorySessionStore()
state = ContinuationState(kind="ltx2.v1", payload={"x": 1})
store.store("sess-1", state)
assert store.snapshot("sess-1") is state
def test_snapshot_missing_returns_none(self):
store = InMemorySessionStore()
assert store.snapshot("missing") is None
def test_store_overwrites_prior_state(self):
store = InMemorySessionStore()
first = ContinuationState(kind="ltx2.v1", payload={"v": 1})
second = ContinuationState(kind="ltx2.v1", payload={"v": 2})
store.store("s", first)
store.store("s", second)
assert store.snapshot("s").payload["v"] == 2
def test_hydrate_assigns_new_session_id(self):
store = InMemorySessionStore()
state = ContinuationState(kind="ltx2.v1", payload={})
sid = store.hydrate(state)
assert sid
assert store.snapshot(sid) is state
def test_hydrate_with_explicit_session_id(self):
store = InMemorySessionStore()
state = ContinuationState(kind="ltx2.v1", payload={})
sid = store.hydrate(state, session_id="pinned-id")
assert sid == "pinned-id"
assert store.snapshot("pinned-id") is state
def test_drop_forgets_session(self):
store = InMemorySessionStore()
store.store("s", ContinuationState(kind="ltx2.v1", payload={}))
store.drop("s")
assert store.snapshot("s") is None
assert "s" not in store
def test_iter_yields_session_ids(self):
store = InMemorySessionStore()
store.store("a", ContinuationState(kind="ltx2.v1", payload={}))
store.store("b", ContinuationState(kind="ltx2.v1", payload={}))
assert sorted(store) == ["a", "b"]
def test_len(self):
store = InMemorySessionStore()
assert len(store) == 0
store.store("x", ContinuationState(kind="ltx2.v1", payload={}))
assert len(store) == 1
def test_concurrent_store_is_safe(self):
"""Smoke-check the lock: 200 parallel stores settle to 200 ids."""
store = InMemorySessionStore()
def write(i: int) -> None:
store.store(
f"s-{i}",
ContinuationState(kind="ltx2.v1", payload={"i": i}),
)
threads = [threading.Thread(target=write, args=(i,)) for i in range(200)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(store) == 200
class TestSnapshotHydrateRoundTrip:
"""Session boundary: snapshot + hydrate preserves the full LTX-2 state."""
def test_end_to_end_ltx2_session_migration(self):
blob_store = InMemoryBlobStore()
sessions = InMemorySessionStore()
typed = LTX2ContinuationState(
segment_index=4,
video_frames=[
np.full((32, 32, 3), i * 5, dtype=np.uint8) for i in range(3)
],
audio_latents=torch.randn(1, 4, 8, 32, dtype=torch.float32),
audio_sample_rate=24000,
audio_conditioning_num_frames=5,
video_position_offset_sec=0.25,
)
envelope = typed.to_continuation_state(blob_store=blob_store)
sessions.store("session-a", envelope)
snapshot = sessions.snapshot("session-a")
assert snapshot is not None
assert snapshot.kind == LTX2_CONTINUATION_KIND
# Simulate a migration: drop the first session, hydrate a new one
# from the snapshot, and reconstruct the typed state.
sessions.drop("session-a")
new_sid = sessions.hydrate(snapshot)
assert new_sid != "session-a"
rebuilt = sessions.snapshot(new_sid)
assert rebuilt is snapshot
restored = LTX2ContinuationState.from_continuation_state(
rebuilt, blob_store=blob_store)
assert restored.segment_index == typed.segment_index
assert restored.audio_sample_rate == typed.audio_sample_rate
torch.testing.assert_close(
restored.audio_latents, typed.audio_latents)
assert len(restored.video_frames) == len(typed.video_frames)
@@ -0,0 +1,99 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the fMP4 encoder.
These tests require ``ffmpeg`` on PATH. Skip when missing so the suite
stays CPU/CI friendly.
"""
from __future__ import annotations
import asyncio
import shutil
import numpy as np
import pytest
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
)
pytestmark = pytest.mark.skipif(
shutil.which("ffmpeg") is None,
reason="ffmpeg not installed",
)
def _frame(width: int, height: int, value: int = 128) -> np.ndarray:
return np.full((height, width, 3), value, dtype=np.uint8)
def test_encoder_emits_init_then_media_chunks():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
chunks: list[FragmentedMP4Chunk] = []
async with enc:
frames = [_frame(64, 64, v) for v in range(4, 28)]
async for chunk in enc.encode(frames):
chunks.append(chunk)
assert len(chunks) > 0
assert chunks[0].kind == "init"
assert all(c.stream_id == enc.stream_id for c in chunks)
assert all(c.segment_idx == 0 for c in chunks)
asyncio.run(run())
def test_encoder_init_chunk_is_fmp4():
"""The first chunk must contain the ``ftyp`` box (fMP4 init segment)."""
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
first_chunk = None
async with enc:
async for chunk in enc.encode([_frame(64, 64, 20)] * 24):
first_chunk = chunk
break
assert first_chunk is not None
assert first_chunk.kind == "init"
# Box header: 4 bytes length, 4 bytes type. "ftyp" should appear
# near the start of the init segment.
assert b"ftyp" in first_chunk.data[:32]
asyncio.run(run())
def test_encoder_rejects_non_ndarray_frames():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
async with enc:
with pytest.raises(TypeError):
async for _ in enc.encode(["not-a-frame"]):
pass
asyncio.run(run())
def test_encoder_rejects_wrong_shape():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
async with enc:
with pytest.raises(ValueError):
async for _ in enc.encode(
[np.zeros((64, 64, 4), dtype=np.uint8)]):
pass
asyncio.run(run())
def test_encoder_close_is_idempotent():
async def run():
enc = FragmentedMP4Encoder(
width=64, height=64, fps=24, segment_idx=0)
await enc.__aenter__()
await enc.close()
await enc.close() # no raise
asyncio.run(run())
+26 -6
View File
@@ -24,6 +24,10 @@ 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", ""),
}))
@@ -66,13 +70,21 @@ 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)
sys.exit(result.returncode)
if result.returncode != 0:
raise RuntimeError(f"Test command failed with exit code {result.returncode}")
# On success, just return — don't call sys.exit()
@app.function(gpu="H100:1",
image=image,
@@ -206,7 +218,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/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
"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"
)
@@ -228,12 +240,20 @@ def run_lora_extraction_tests():
timeout=1800,
secrets=[
modal.Secret.from_dict(
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
{"HF_API_KEY": os.environ.get("HF_API_KEY", ""),
"HF_REPO_ID": "FastVideo/performance-tracking"})
],
volumes={"/root/data": model_vol})
volumes={
"/root/data": model_vol,
})
def run_performance_tests():
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/performance -vs"
"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"
)
+1 -1
View File
@@ -480,10 +480,10 @@ def _prepare_ssim_workspace(
{checkout_command}
rm -rf fastvideo/tests/ssim/reference_videos
git_retry git submodule update --init --recursive
uv pip install -e .[test]
cd fastvideo-kernel
./build.sh
cd ..
uv pip install -e .[test]
uv pip install git+https://github.com/microsoft/MoGe.git
export HF_HOME='/root/data/.cache'
hf auth login --token "$HF_API_KEY"
@@ -0,0 +1,320 @@
# 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())

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