Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b3a9874fc8 | ||
|
|
d657cbbf17 | ||
|
|
9801037c3d | ||
|
|
74d09b0efd | ||
|
|
38dc8820ac | ||
|
|
c77a76c6af | ||
|
|
d14d5aadea | ||
|
|
4c915b7742 | ||
|
|
48534ef4de | ||
|
|
1116f514be | ||
|
|
d451e61749 | ||
|
|
9a8bbe18fa | ||
|
|
ea25441ef0 | ||
|
|
48957fcde1 | ||
|
|
7b872cc41e | ||
|
|
3ff4a8d2d2 | ||
|
|
9343d4cdf4 | ||
|
|
66fb3d1e79 | ||
|
|
37418946c8 | ||
|
|
95fd29e0cb | ||
|
|
aca850cef2 | ||
|
|
1c79779956 | ||
|
|
eee03527ed | ||
|
|
1eb8541094 | ||
|
|
e17cd2633c | ||
|
|
e0dc5f2b0c | ||
|
|
69c214d13a | ||
|
|
0341481aa7 | ||
|
|
70ee5d230c | ||
|
|
d1c3fdd187 | ||
|
|
980e8d933e | ||
|
|
24ced500f5 |
Executable
+96
@@ -0,0 +1,96 @@
|
||||
#!/usr/bin/env bash
|
||||
# Sync .agents/skills/ into .claude/skills/ via per-skill symlinks.
|
||||
#
|
||||
# Why: Claude Code only scans .claude/skills/ and ~/.claude/skills/ for
|
||||
# user-invocable skills (no skillsPath config exists — see
|
||||
# https://code.claude.com/docs/en/skills.md). This repo's skills live
|
||||
# in .agents/skills/ so they travel with the repo and stay under git.
|
||||
# Run this once after cloning (or after adding/removing a skill) to
|
||||
# expose them to Claude Code without maintaining a parallel tree.
|
||||
#
|
||||
# Usage:
|
||||
# .agents/scripts/sync-skills.sh
|
||||
#
|
||||
# Idempotent and safe to re-run. Prunes stale symlinks whose source
|
||||
# has been removed from .agents/skills/. Leaves hand-written
|
||||
# .claude/skills/<name>/ directories untouched (only symlinks are
|
||||
# managed).
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(git -C "$(dirname "$0")" rev-parse --show-toplevel)"
|
||||
SRC_DIR="$REPO_ROOT/.agents/skills"
|
||||
DST_DIR="$REPO_ROOT/.claude/skills"
|
||||
|
||||
if [[ ! -d "$SRC_DIR" ]]; then
|
||||
echo "Error: $SRC_DIR does not exist." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir -p "$DST_DIR"
|
||||
|
||||
linked=0
|
||||
unchanged=0
|
||||
skipped=0
|
||||
pruned=0
|
||||
|
||||
link_skill() {
|
||||
local name="$1"
|
||||
local src="$SRC_DIR/$name"
|
||||
local dst="$DST_DIR/$name"
|
||||
# Relative target keeps symlinks portable across clones.
|
||||
local rel="../../.agents/skills/$name"
|
||||
|
||||
if [[ -L "$dst" ]]; then
|
||||
if [[ "$(readlink "$dst")" == "$rel" ]]; then
|
||||
unchanged=$((unchanged + 1))
|
||||
return
|
||||
fi
|
||||
rm "$dst"
|
||||
elif [[ -e "$dst" ]]; then
|
||||
echo "Skipped (not a symlink): .claude/skills/$name" >&2
|
||||
skipped=$((skipped + 1))
|
||||
return
|
||||
fi
|
||||
|
||||
ln -s "$rel" "$dst"
|
||||
echo "Linked: .claude/skills/$name -> $rel"
|
||||
linked=$((linked + 1))
|
||||
}
|
||||
|
||||
prune_stale() {
|
||||
local link="$1"
|
||||
local target
|
||||
target="$(readlink "$link")"
|
||||
case "$target" in
|
||||
../../.agents/skills/*) ;;
|
||||
*) return ;;
|
||||
esac
|
||||
local name="${target##*/}"
|
||||
if [[ ! -d "$SRC_DIR/$name" ]]; then
|
||||
rm "$link"
|
||||
echo "Pruned stale: .claude/skills/$(basename "$link")"
|
||||
pruned=$((pruned + 1))
|
||||
fi
|
||||
}
|
||||
|
||||
for src in "$SRC_DIR"/*/; do
|
||||
[[ -d "$src" ]] || continue
|
||||
name="$(basename "$src")"
|
||||
# Only treat directories that actually contain a SKILL.md as skills.
|
||||
[[ -f "$src/SKILL.md" ]] || continue
|
||||
link_skill "$name"
|
||||
done
|
||||
|
||||
shopt -s nullglob
|
||||
for link in "$DST_DIR"/*; do
|
||||
[[ -L "$link" ]] || continue
|
||||
prune_stale "$link"
|
||||
done
|
||||
shopt -u nullglob
|
||||
|
||||
printf "\nSummary: %d linked, %d unchanged, %d pruned" "$linked" "$unchanged" "$pruned"
|
||||
if [[ "$skipped" -gt 0 ]]; then
|
||||
printf ", %d skipped (non-symlink collision)" "$skipped"
|
||||
fi
|
||||
printf "\n"
|
||||
@@ -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": {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
# Streaming WebSocket Server Contract
|
||||
|
||||
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
|
||||
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
|
||||
document is the authoritative spec for the message catalogue and the
|
||||
session state machine. Any change to either must update this document
|
||||
in the same PR that touches `protocol.py` or `session.py`.
|
||||
|
||||
## Endpoint
|
||||
|
||||
| Path | Protocol | Purpose |
|
||||
|---|---|---|
|
||||
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
|
||||
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
|
||||
|
||||
The server is launched by `fastvideo serve --config <serve.yaml>` when
|
||||
the config carries a `streaming:` block. Without that block the same CLI
|
||||
launches the OpenAI stateless HTTP server instead.
|
||||
|
||||
## Connection lifecycle
|
||||
|
||||
Every WebSocket connection holds exactly one `Session`. Sessions move
|
||||
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
|
||||
|
||||
```
|
||||
┌──────────────┐
|
||||
│ INITIALIZING │ ← WebSocket accepted, before init frame
|
||||
└──────┬───────┘
|
||||
│ session_init_v2 received
|
||||
┌──────────────┼──────────────┐
|
||||
▼ ▼ ▼
|
||||
QUEUED GPU_BINDING REJECTED
|
||||
│ │ ↑
|
||||
│ slot ready │ │ max-sessions hit
|
||||
▼ ▼ │ or invalid init
|
||||
┌────────┐ │
|
||||
│ ACTIVE │ ────────┘
|
||||
└────┬───┘
|
||||
segment loop │
|
||||
│
|
||||
┌───────────┼───────────┐
|
||||
▼ ▼ ▼
|
||||
COMPLETE ERROR TIMEOUT
|
||||
(clean leave) (any failure) (idle / segment_cap reached)
|
||||
```
|
||||
|
||||
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
|
||||
no transitions out. The transition matrix is enforced in
|
||||
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
|
||||
|
||||
`SessionManager` enforces the per-process budgets pulled from
|
||||
`StreamingConfig`:
|
||||
|
||||
- `session_timeout_seconds` — idle reaper drops sessions that haven't
|
||||
advanced; non-terminal sessions transition to `TIMEOUT`.
|
||||
- `generation_segment_cap` — a session that hits the cap transitions to
|
||||
`COMPLETE` after the last segment ships.
|
||||
|
||||
## Message catalogue
|
||||
|
||||
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
|
||||
`protocol.py` are the source of truth; this table is the human-readable
|
||||
view.
|
||||
|
||||
### Client → server
|
||||
|
||||
| `type` | Required fields | Purpose |
|
||||
|---|---|---|
|
||||
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
|
||||
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
|
||||
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
|
||||
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
|
||||
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
|
||||
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
|
||||
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
|
||||
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
|
||||
|
||||
The opening frame must be `session_init_v2`. Any other first frame is
|
||||
rejected with an `error` (code `invalid_message`) and the WebSocket is
|
||||
closed.
|
||||
|
||||
### Server → client
|
||||
|
||||
| `type` | Carries | When emitted |
|
||||
|---|---|---|
|
||||
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
|
||||
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
|
||||
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
|
||||
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
|
||||
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
|
||||
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
|
||||
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
|
||||
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
|
||||
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
|
||||
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
|
||||
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
|
||||
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
|
||||
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
|
||||
|
||||
## Continuation state
|
||||
|
||||
The session optionally accepts a `continuation_state` dict inside the
|
||||
opening `session_init_v2` frame. When present, the server hydrates it
|
||||
into a `ContinuationState(kind, payload)` envelope and feeds it as the
|
||||
`request.state` on the first segment's `GenerationRequest` — letting a
|
||||
client resume after a disconnect, migrate sessions across processes,
|
||||
or replay a prior session.
|
||||
|
||||
After every segment, if the runtime returns a fresh state, the server
|
||||
persists it to the `SessionStore` so a `snapshot_state` request can
|
||||
export it. The store and serialization contracts live with the model
|
||||
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
|
||||
|
||||
## Example flow
|
||||
|
||||
```
|
||||
client server
|
||||
────── ──────
|
||||
WS /v1/stream ─────── connect ─────────────────────────►
|
||||
◄────── (accept)
|
||||
|
||||
{"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox in snow", "the fox jumps"],
|
||||
"initial_image": {...},
|
||||
"stream_mode": "av_fmp4"} ─────────────────────────────►
|
||||
|
||||
(validate, queue, bind)
|
||||
◄──── {"type": "queue_status",
|
||||
"position": 0, "queue_depth": 0}
|
||||
◄──── {"type": "gpu_assigned",
|
||||
"gpu_id": 0, "model_id": "..."}
|
||||
◄──── {"type": "ltx2_stream_start", ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "a fox in snow",
|
||||
"source": "curated"} ───────────────────────────────────►
|
||||
(run pipeline)
|
||||
◄──── {"type": "ltx2_segment_start",
|
||||
"segment_idx": 1, ...}
|
||||
◄──── {"type": "step_complete",
|
||||
"segment_idx": 1, "timings": {...}}
|
||||
◄──── {"type": "media_init",
|
||||
"segment_idx": 1,
|
||||
"mime": "video/mp4", ...}
|
||||
◄──── <binary fMP4 init segment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── {"type": "media_segment_complete",
|
||||
"segment_idx": 1, "chunks": 12}
|
||||
◄──── {"type": "ltx2_segment_complete",
|
||||
"segment_idx": 1, ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "the fox jumps"} ─────────────────────────────►
|
||||
(segment 2 …)
|
||||
|
||||
{"type": "snapshot_state"} ──────────────────────────────►
|
||||
◄──── {"type": "continuation_state_snapshot",
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, ...}}
|
||||
|
||||
(close) ──────────────────────────────────────────────────►
|
||||
(session → COMPLETE)
|
||||
```
|
||||
|
||||
## Backward / forward compatibility
|
||||
|
||||
- Adding a new client message: append a Pydantic model to `protocol.py`
|
||||
with a unique `type`; add the discriminator entry to `ClientMessage`;
|
||||
add a row to the table above. Old clients that don't send the new
|
||||
message remain compatible.
|
||||
- Adding a new server message: emit only when a new feature flag is
|
||||
enabled (or always emit, since clients ignore unknown types).
|
||||
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
|
||||
→ `session_init_v3`) and accept both for one release cycle. Never
|
||||
silently change field semantics under the same `type`.
|
||||
@@ -86,3 +86,25 @@ sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
|
||||
- Learning rate: 2e-5
|
||||
- Training steps: 3000 (~12 hours)
|
||||
- HSDP shard dim: 1
|
||||
|
||||
## 🧭 Note on `real_score_guidance_scale`
|
||||
|
||||
The teacher CFG used inside the DMD loss follows the DMD2 reference
|
||||
implementation and uses the parameterization
|
||||
|
||||
```
|
||||
x = x_cond + w * (x_cond - x_uncond)
|
||||
```
|
||||
|
||||
rather than the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)`. The
|
||||
two are mathematically equivalent up to a constant offset:
|
||||
|
||||
| `real_score_guidance_scale` (`w`) | Equivalent standard CFG (`w + 1`) | Output |
|
||||
|-----------------------------------|-----------------------------------|-----------------------|
|
||||
| `-1` | `0` | unconditional |
|
||||
| `0` | `1` | conditional |
|
||||
| `3.5` (default) | `4.5` | strong guidance |
|
||||
|
||||
So `real_score_guidance_scale` should be read as the **extra** guidance
|
||||
strength added on top of the conditional prediction. When porting values
|
||||
from a paper that uses the Ho & Salimans form, subtract 1.
|
||||
|
||||
@@ -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
|
||||
@@ -61,7 +61,16 @@ has_cmake_arg() {
|
||||
}
|
||||
|
||||
detect_with_torch() {
|
||||
uv run --active --no-project python -c "import torch
|
||||
# Prefer the active venv's python directly over `uv run --active --no-project`,
|
||||
# which on some uv versions provisions its own interpreter and misses packages
|
||||
# installed into VIRTUAL_ENV.
|
||||
local py
|
||||
if [[ -n "${VIRTUAL_ENV:-}" && -x "${VIRTUAL_ENV}/bin/python" ]]; then
|
||||
py="${VIRTUAL_ENV}/bin/python"
|
||||
else
|
||||
py="$(command -v python3 || command -v python)"
|
||||
fi
|
||||
"${py}" -c "import torch
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError('torch.cuda.is_available() is false')
|
||||
mj, mn = torch.cuda.get_device_capability(0)
|
||||
|
||||
@@ -5,6 +5,11 @@ from fastvideo_kernel.ops import (
|
||||
video_sparse_attn,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn import (
|
||||
block_sparse_attn,
|
||||
block_sparse_attn_from_indices,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.vmoba import (
|
||||
moba_attn_varlen,
|
||||
process_moba_input,
|
||||
@@ -22,6 +27,8 @@ from fastvideo_kernel.turbodiffusion_ops import (
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"block_sparse_attn",
|
||||
"block_sparse_attn_from_indices",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
@@ -6,6 +8,11 @@ from typing import Tuple
|
||||
import torch
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Backend selection helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_sm90_ops():
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
|
||||
@@ -25,38 +32,66 @@ def _is_sm90() -> bool:
|
||||
|
||||
|
||||
def _force_triton() -> bool:
|
||||
# Force Triton even on SM90 and even if the compiled extension is available.
|
||||
# Useful for CI / debugging / parity testing.
|
||||
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
|
||||
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Preferred map->index conversion used by the wrapper.
|
||||
# ---------------------------------------------------------------------------
|
||||
# Index helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
This wrapper **requires** the Triton implementation.
|
||||
If Triton (or the Triton map_to_index module) is not available, it raises.
|
||||
"""
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Compact a bool block_map to (q2k_idx, q2k_num). Legacy path only."""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
if block_map.dim() != 4:
|
||||
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
|
||||
raise ValueError(
|
||||
f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), "
|
||||
f"got shape={tuple(block_map.shape)}"
|
||||
)
|
||||
if block_map.dtype != torch.bool:
|
||||
block_map = block_map.to(torch.bool)
|
||||
|
||||
if not block_map.is_cuda:
|
||||
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
|
||||
raise RuntimeError(
|
||||
"block_map must be a CUDA tensor (Triton map_to_index required)."
|
||||
)
|
||||
|
||||
try:
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
|
||||
except Exception as e:
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index
|
||||
except Exception as e: # pragma: no cover - environment issue
|
||||
raise ImportError(
|
||||
"Triton map_to_index is required but not available. "
|
||||
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
|
||||
"Ensure Triton is installed and "
|
||||
"fastvideo_kernel.triton_kernels.index is importable."
|
||||
) from e
|
||||
return triton_map_to_index(block_map)
|
||||
|
||||
|
||||
def _invert_indices_for_backward(
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from fastvideo_kernel.triton_kernels.index import invert_indices
|
||||
return invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
|
||||
|
||||
|
||||
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
|
||||
"""Return `t` as a contiguous int32 tensor, raising a clear error on CPU input."""
|
||||
if not t.is_cuda:
|
||||
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
|
||||
if t.dtype != torch.int32:
|
||||
t = t.to(torch.int32)
|
||||
if not t.is_contiguous():
|
||||
t = t.contiguous()
|
||||
return t
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Triton backend custom ops (index-native)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_triton",
|
||||
mutates_args=(),
|
||||
@@ -66,34 +101,40 @@ def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
|
||||
triton_block_sparse_attn_forward,
|
||||
)
|
||||
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
o, M = triton_block_sparse_attn_forward(
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
return o, M
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
M = torch.empty(
|
||||
(q.shape[0], q.shape[1], q.shape[2]),
|
||||
device=q.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
return o, M
|
||||
|
||||
|
||||
@@ -109,20 +150,32 @@ def block_sparse_attn_backward_triton(
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output = grad_output.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
|
||||
triton_block_sparse_attn_backward,
|
||||
)
|
||||
|
||||
num_kv_blocks = int(variable_block_sizes.numel())
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(
|
||||
q2k_idx, q2k_num, num_kv_blocks
|
||||
)
|
||||
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
|
||||
# o/M are kernel outputs so are already contiguous.
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(
|
||||
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
|
||||
grad_output.contiguous(),
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
o,
|
||||
M,
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
return dq, dk, dv
|
||||
|
||||
@@ -135,7 +188,8 @@ def _block_sparse_attn_backward_triton_fake(
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q)
|
||||
@@ -144,19 +198,28 @@ def _block_sparse_attn_backward_triton_fake(
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _backward_triton(ctx, grad_o, grad_M):
|
||||
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _setup_context_triton(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
|
||||
o, M = output
|
||||
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
|
||||
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
|
||||
def _backward_triton(ctx, grad_o, grad_M):
|
||||
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(
|
||||
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None, None
|
||||
|
||||
|
||||
block_sparse_attn_triton.register_autograd(
|
||||
_backward_triton, setup_context=_setup_context_triton
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SM90 backend custom ops (index-native)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
@@ -168,21 +231,21 @@ def block_sparse_attn_sm90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
block_sparse_fwd, _ = _get_sm90_ops()
|
||||
if block_sparse_fwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
|
||||
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
o_padded, lse_padded = block_sparse_fwd(
|
||||
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
|
||||
q_padded.contiguous(),
|
||||
k_padded.contiguous(),
|
||||
v_padded.contiguous(),
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
return o_padded, lse_padded
|
||||
|
||||
@@ -192,11 +255,16 @@ def _block_sparse_attn_sm90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q_padded)
|
||||
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
|
||||
lse = torch.empty(
|
||||
(q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1),
|
||||
device=q_padded.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
return o, lse
|
||||
|
||||
|
||||
@@ -212,30 +280,34 @@ def block_sparse_attn_backward_sm90(
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
_, block_sparse_bwd = _get_sm90_ops()
|
||||
if block_sparse_bwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
|
||||
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
num_kv_blocks = int(variable_block_sizes.numel())
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(
|
||||
q2k_idx, q2k_num, num_kv_blocks
|
||||
)
|
||||
|
||||
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
|
||||
dq, dk, dv = block_sparse_bwd(
|
||||
q_padded,
|
||||
k_padded,
|
||||
v_padded,
|
||||
q_padded.contiguous(),
|
||||
k_padded.contiguous(),
|
||||
v_padded.contiguous(),
|
||||
o_padded,
|
||||
lse_padded,
|
||||
grad_output_padded,
|
||||
grad_output_padded.contiguous(),
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes.int(),
|
||||
variable_block_sizes,
|
||||
)
|
||||
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
|
||||
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
|
||||
# C++ kernel returns fp32 grads; cast back to the input dtype.
|
||||
out_dtype = grad_output_padded.dtype
|
||||
return dq.to(out_dtype), dk.to(out_dtype), dv.to(out_dtype)
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
|
||||
@@ -246,7 +318,8 @@ def _block_sparse_attn_backward_sm90_fake(
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q_padded)
|
||||
@@ -255,21 +328,57 @@ def _block_sparse_attn_backward_sm90_fake(
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _backward_sm90(ctx, grad_o, grad_lse):
|
||||
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_sm90(
|
||||
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _setup_context_sm90(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
|
||||
o, lse = output
|
||||
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
|
||||
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
|
||||
def _backward_sm90(ctx, grad_o, grad_lse):
|
||||
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_sm90(
|
||||
grad_o, q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None, None
|
||||
|
||||
|
||||
block_sparse_attn_sm90.register_autograd(
|
||||
_backward_sm90, setup_context=_setup_context_sm90
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def block_sparse_attn_from_indices(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Block-sparse attention with autograd, taking compact per-row KV indices."""
|
||||
# Normalize index tensors once at the public boundary so the custom ops
|
||||
# and their fakes can assume int32/contiguous. No-op on well-formed input.
|
||||
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
|
||||
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
|
||||
variable_block_sizes = _as_int32_contig(variable_block_sizes, "variable_block_sizes")
|
||||
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
use_sm90 = (
|
||||
(not _force_triton())
|
||||
and _is_sm90()
|
||||
and block_sparse_fwd is not None
|
||||
and block_sparse_bwd is not None
|
||||
)
|
||||
if use_sm90:
|
||||
return block_sparse_attn_sm90(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
|
||||
# to a multiple of the block size (64 tokens).
|
||||
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
def block_sparse_attn(
|
||||
@@ -279,16 +388,8 @@ def block_sparse_attn(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Unified block-sparse attention op with autograd support.
|
||||
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
|
||||
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
|
||||
"""
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
|
||||
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
|
||||
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
|
||||
# to a multiple of the block size (64 tokens).
|
||||
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
"""Bool-mask compat wrapper; prefer block_sparse_attn_from_indices."""
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
return block_sparse_attn_from_indices(
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import math
|
||||
import torch
|
||||
from .block_sparse_attn import block_sparse_attn
|
||||
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
|
||||
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
|
||||
|
||||
# Try to load the C++ extension
|
||||
@@ -125,13 +125,18 @@ def video_sparse_attn(
|
||||
out_c = out_c.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch, heads, q_seq_len, dim)
|
||||
|
||||
# Sparse branch
|
||||
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
|
||||
topk_idx = torch.topk(scores, topk, dim=-1).indices
|
||||
mask = torch.zeros_like(scores,
|
||||
dtype=torch.bool).scatter_(-1, topk_idx, True)
|
||||
|
||||
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
q2k_idx = topk_idx.to(torch.int32).contiguous()
|
||||
q2k_num = torch.full(
|
||||
(batch, heads, q_num_blocks),
|
||||
topk,
|
||||
dtype=torch.int32,
|
||||
device=q.device,
|
||||
)
|
||||
out_s = block_sparse_attn_from_indices(
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes
|
||||
)[0]
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
## pytorch sdpa version of block sparse ##
|
||||
from typing import Tuple
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
@@ -153,3 +154,114 @@ def map_to_index(block_map: torch.Tensor):
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _invert_indices_kernel(
|
||||
q2k_idx_ptr,
|
||||
q2k_num_ptr,
|
||||
k2q_idx_ptr,
|
||||
k2q_num_ptr,
|
||||
q2k_idx_b, q2k_idx_h, q2k_idx_q, q2k_idx_k,
|
||||
q2k_num_b, q2k_num_h, q2k_num_q,
|
||||
k2q_idx_b, k2q_idx_h, k2q_idx_k, k2q_idx_q,
|
||||
k2q_num_b, k2q_num_h, k2q_num_k,
|
||||
MAX_KV_PER_Q: tl.constexpr,
|
||||
):
|
||||
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
|
||||
pid_b = tl.program_id(0)
|
||||
pid_h = tl.program_id(1)
|
||||
pid_q = tl.program_id(2)
|
||||
|
||||
n = tl.load(
|
||||
q2k_num_ptr
|
||||
+ pid_b * q2k_num_b
|
||||
+ pid_h * q2k_num_h
|
||||
+ pid_q * q2k_num_q
|
||||
)
|
||||
|
||||
q2k_row = (
|
||||
q2k_idx_ptr
|
||||
+ pid_b * q2k_idx_b
|
||||
+ pid_h * q2k_idx_h
|
||||
+ pid_q * q2k_idx_q
|
||||
)
|
||||
|
||||
for i in tl.range(0, MAX_KV_PER_Q):
|
||||
if i < n:
|
||||
kv = tl.load(q2k_row + i * q2k_idx_k)
|
||||
count_ptr = (
|
||||
k2q_num_ptr
|
||||
+ pid_b * k2q_num_b
|
||||
+ pid_h * k2q_num_h
|
||||
+ kv * k2q_num_k
|
||||
)
|
||||
pos = tl.atomic_add(count_ptr, 1)
|
||||
tl.store(
|
||||
k2q_idx_ptr
|
||||
+ pid_b * k2q_idx_b
|
||||
+ pid_h * k2q_idx_h
|
||||
+ kv * k2q_idx_k
|
||||
+ pos * k2q_idx_q,
|
||||
pid_q,
|
||||
)
|
||||
|
||||
|
||||
def invert_indices(
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
|
||||
if q2k_idx.dim() != 4:
|
||||
raise ValueError(
|
||||
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
|
||||
)
|
||||
if q2k_num.dim() != 3:
|
||||
raise ValueError(
|
||||
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
|
||||
)
|
||||
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
|
||||
raise RuntimeError("invert_indices requires CUDA tensors.")
|
||||
|
||||
B, H, Nq, Mk = q2k_idx.shape
|
||||
if q2k_num.shape != (B, H, Nq):
|
||||
raise ValueError(
|
||||
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
|
||||
f"[B, H, Nq] = {(B, H, Nq)}"
|
||||
)
|
||||
|
||||
q2k_idx = q2k_idx.contiguous()
|
||||
q2k_num = q2k_num.contiguous()
|
||||
if q2k_idx.dtype != torch.int32:
|
||||
q2k_idx = q2k_idx.to(torch.int32)
|
||||
if q2k_num.dtype != torch.int32:
|
||||
q2k_num = q2k_num.to(torch.int32)
|
||||
|
||||
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
|
||||
# `Nq` is a tight upper bound on the compacted K->Q slots.
|
||||
k2q_idx = torch.empty(
|
||||
(B, H, num_kv_blocks, Nq),
|
||||
dtype=torch.int32,
|
||||
device=q2k_idx.device,
|
||||
)
|
||||
k2q_num = torch.zeros(
|
||||
(B, H, num_kv_blocks),
|
||||
dtype=torch.int32,
|
||||
device=q2k_idx.device,
|
||||
)
|
||||
|
||||
grid = (B, H, Nq)
|
||||
_invert_indices_kernel[grid](
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
q2k_idx.stride(0), q2k_idx.stride(1), q2k_idx.stride(2), q2k_idx.stride(3),
|
||||
q2k_num.stride(0), q2k_num.stride(1), q2k_num.stride(2),
|
||||
k2q_idx.stride(0), k2q_idx.stride(1), k2q_idx.stride(2), k2q_idx.stride(3),
|
||||
k2q_num.stride(0), k2q_num.stride(1), k2q_num.stride(2),
|
||||
MAX_KV_PER_Q=Mk,
|
||||
)
|
||||
|
||||
return k2q_idx, k2q_num
|
||||
|
||||
+118
-9
@@ -16,6 +16,8 @@ from fastvideo.api.request_metadata import (
|
||||
reset_tracking_roots,
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
@@ -25,6 +27,10 @@ from fastvideo.api.schema import (
|
||||
)
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
from fastvideo.utils import shallow_asdict
|
||||
|
||||
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
|
||||
@@ -38,6 +44,10 @@ _LEGACY_REQUEST_ALIASES = {
|
||||
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
|
||||
"embedded_cfg_scale",
|
||||
})
|
||||
# torch.compile kwargs that map to first-class CompileConfig fields.
|
||||
_COMPILE_TYPED_KEYS = ("backend", "fullgraph", "mode", "dynamic")
|
||||
# LTX-2 refine flat kwargs (init + per-request) known to FastVideoArgs.
|
||||
_LTX2_REFINE_FLAT_KEYS = (refine_preset_override_fields() | refine_stage_override_fields())
|
||||
|
||||
|
||||
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
|
||||
@@ -80,6 +90,8 @@ def legacy_from_pretrained_to_config(
|
||||
components: dict[str, Any] = {}
|
||||
quantization: dict[str, Any] = {}
|
||||
experimental: dict[str, Any] = {}
|
||||
preset_overrides: dict[str, Any] = {}
|
||||
preset_refine: dict[str, Any] = {}
|
||||
|
||||
for key, value in kwargs.items():
|
||||
if key == "revision":
|
||||
@@ -106,8 +118,33 @@ def legacy_from_pretrained_to_config(
|
||||
offload["pin_cpu_memory"] = value
|
||||
elif key == "enable_torch_compile":
|
||||
compile_config["enabled"] = value
|
||||
elif key == "enable_torch_compile_text_encoder":
|
||||
compile_config["text_encoder_enabled"] = value
|
||||
elif key == "torch_compile_kwargs":
|
||||
compile_config["kwargs"] = deepcopy(value)
|
||||
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
|
||||
for first_class in _COMPILE_TYPED_KEYS:
|
||||
if first_class in remaining:
|
||||
compile_config[first_class] = remaining.pop(first_class)
|
||||
if remaining:
|
||||
compile_config["extras"] = remaining
|
||||
elif key == "ltx2_vae_tiling":
|
||||
pipeline["vae_tiling"] = value
|
||||
elif key == "config_model_path":
|
||||
components["config_root"] = value
|
||||
elif key == "ltx2_refine_enabled":
|
||||
preset_refine["enabled"] = value
|
||||
elif key == "ltx2_refine_upsampler_path":
|
||||
# Empty string means "no upsampler"; keep typed None.
|
||||
components["upsampler_weights"] = value or None
|
||||
elif key == "ltx2_refine_lora_path":
|
||||
# Empty string means "no refine LoRA"; keep typed None.
|
||||
components["lora_path"] = value or None
|
||||
elif key == "ltx2_refine_add_noise":
|
||||
preset_refine["add_noise"] = value
|
||||
elif key == "ltx2_refine_num_inference_steps":
|
||||
preset_refine["num_inference_steps"] = value
|
||||
elif key == "ltx2_refine_guidance_scale":
|
||||
preset_refine["guidance_scale"] = value
|
||||
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
|
||||
engine[key] = value
|
||||
elif key == "override_text_encoder_quant":
|
||||
@@ -147,6 +184,10 @@ def legacy_from_pretrained_to_config(
|
||||
|
||||
if components:
|
||||
pipeline["components"] = components
|
||||
if preset_refine:
|
||||
preset_overrides["refine"] = preset_refine
|
||||
if preset_overrides:
|
||||
pipeline["preset_overrides"] = preset_overrides
|
||||
if experimental:
|
||||
pipeline["experimental"] = experimental
|
||||
if pipeline:
|
||||
@@ -162,12 +203,8 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
unsupported.append("pipeline.preset")
|
||||
if normalized.pipeline.preset_version is not None:
|
||||
unsupported.append("pipeline.preset_version")
|
||||
if normalized.pipeline.components.config_root is not None:
|
||||
unsupported.append("pipeline.components.config_root")
|
||||
if normalized.pipeline.components.vae_weights is not None:
|
||||
unsupported.append("pipeline.components.vae_weights")
|
||||
if normalized.pipeline.components.upsampler_weights is not None:
|
||||
unsupported.append("pipeline.components.upsampler_weights")
|
||||
if unsupported:
|
||||
joined = ", ".join(unsupported)
|
||||
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
|
||||
@@ -191,13 +228,21 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
"vae_cpu_offload": engine.offload.vae,
|
||||
"pin_cpu_memory": engine.offload.pin_cpu_memory,
|
||||
"enable_torch_compile": engine.compile.enabled,
|
||||
"torch_compile_kwargs": deepcopy(engine.compile.kwargs),
|
||||
"torch_compile_kwargs": _compile_config_to_torch_kwargs(engine.compile),
|
||||
"enable_stage_verification": engine.enable_stage_verification,
|
||||
"use_fsdp_inference": engine.use_fsdp_inference,
|
||||
"disable_autocast": engine.disable_autocast,
|
||||
}
|
||||
if normalized.pipeline.workload_type is not None:
|
||||
kwargs["workload_type"] = normalized.pipeline.workload_type
|
||||
if normalized.pipeline.vae_tiling is not None:
|
||||
kwargs["ltx2_vae_tiling"] = normalized.pipeline.vae_tiling
|
||||
if engine.compile.text_encoder_enabled is not None:
|
||||
# ``FastVideoArgs.from_kwargs`` filters to declared fields, so
|
||||
# this is a no-op on the current legacy path. Emit anyway so the
|
||||
# realtime runtime (PR 7.6) — which reads from the kwargs dict
|
||||
# before FastVideoArgs filtering — can pick it up once wired.
|
||||
kwargs["enable_torch_compile_text_encoder"] = (engine.compile.text_encoder_enabled)
|
||||
|
||||
quantization = engine.quantization
|
||||
if quantization is not None and quantization.text_encoder_quant is not None:
|
||||
@@ -220,8 +265,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
kwargs["init_weights_from_safetensors"] = components.transformer_weights
|
||||
if components.transformer_2_weights is not None:
|
||||
kwargs["init_weights_from_safetensors_2"] = components.transformer_2_weights
|
||||
if components.config_root is not None:
|
||||
kwargs["config_model_path"] = components.config_root
|
||||
if components.upsampler_weights is not None:
|
||||
kwargs["ltx2_refine_upsampler_path"] = components.upsampler_weights
|
||||
|
||||
kwargs.update(deepcopy(normalized.pipeline.preset_overrides))
|
||||
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
|
||||
refine = preset_overrides.pop("refine", None)
|
||||
if isinstance(refine, Mapping):
|
||||
for key in _LTX2_REFINE_FLAT_KEYS:
|
||||
if key in refine:
|
||||
kwargs[f"ltx2_refine_{key}"] = refine[key]
|
||||
kwargs.update(preset_overrides)
|
||||
kwargs.update(deepcopy(normalized.pipeline.experimental))
|
||||
return FastVideoArgs.from_kwargs(**kwargs)
|
||||
|
||||
@@ -271,10 +326,13 @@ def request_to_sampling_param(
|
||||
) -> SamplingParam:
|
||||
if request.plan is not None:
|
||||
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
|
||||
if request.state is not None:
|
||||
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
if request.state is not None:
|
||||
_validate_continuation_state(request.state)
|
||||
sampling_param.continuation_state = request.state
|
||||
if request.output.return_state:
|
||||
sampling_param.return_continuation_state = True
|
||||
updates = explicit_request_updates(request)
|
||||
|
||||
for key, value in updates.items():
|
||||
@@ -316,6 +374,25 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
|
||||
return isinstance(raw.get("generator"), Mapping)
|
||||
|
||||
|
||||
def _compile_config_to_torch_kwargs(compile_config: CompileConfig, ) -> dict[str, Any]:
|
||||
"""Flatten typed ``CompileConfig`` back to a ``torch_compile_kwargs``
|
||||
dict that the legacy ``FastVideoArgs`` path still expects.
|
||||
|
||||
Typed first-class fields (:attr:`backend`, :attr:`fullgraph`,
|
||||
:attr:`mode`, :attr:`dynamic`) are only emitted when the user set
|
||||
them explicitly (non-``None``). ``extras`` is merged on top for any
|
||||
uncommon kwargs.
|
||||
"""
|
||||
out: dict[str, Any] = {}
|
||||
for key in _COMPILE_TYPED_KEYS:
|
||||
value = getattr(compile_config, key)
|
||||
if value is not None:
|
||||
out[key] = value
|
||||
if compile_config.extras:
|
||||
out.update(deepcopy(compile_config.extras))
|
||||
return out
|
||||
|
||||
|
||||
def _sampling_param_to_request_raw(sampling_param: SamplingParam | None, ) -> dict[str, Any]:
|
||||
if sampling_param is None:
|
||||
return {}
|
||||
@@ -476,6 +553,37 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
|
||||
|
||||
_SCHEMA_DEFAULT_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
|
||||
|
||||
_KNOWN_CONTINUATION_KINDS: set[str] = set()
|
||||
|
||||
|
||||
def register_continuation_kind(kind: str) -> None:
|
||||
"""Register a :class:`ContinuationState.kind` as recognized.
|
||||
|
||||
PR 7 wires the envelope through; per-kind payload deserializers live
|
||||
with each model family (e.g. ``fastvideo.pipelines.basic.ltx2.
|
||||
continuation.LTX2ContinuationState``). The registry lets the
|
||||
public-API compat layer validate the kind early, before the state
|
||||
reaches the pipeline.
|
||||
"""
|
||||
if not isinstance(kind, str) or not kind:
|
||||
raise ValueError("ContinuationState kind must be a non-empty string")
|
||||
_KNOWN_CONTINUATION_KINDS.add(kind)
|
||||
|
||||
|
||||
def _validate_continuation_state(state: ContinuationState) -> None:
|
||||
if not isinstance(state.kind, str) or not state.kind:
|
||||
raise ValueError("GenerationRequest.state.kind must be a non-empty string; got "
|
||||
f"{state.kind!r}")
|
||||
if not isinstance(state.payload, Mapping):
|
||||
raise ValueError(f"GenerationRequest.state.payload must be a mapping; got "
|
||||
f"{type(state.payload).__name__}")
|
||||
if state.kind not in _KNOWN_CONTINUATION_KINDS:
|
||||
known = sorted(_KNOWN_CONTINUATION_KINDS)
|
||||
raise ValueError(f"Unknown ContinuationState kind {state.kind!r}; registered "
|
||||
f"kinds: {known}. Import the model family that owns this kind "
|
||||
"(e.g. `import fastvideo.pipelines.basic.ltx2.continuation`) "
|
||||
"to register it, or drop the state field.")
|
||||
|
||||
|
||||
def _fan_out_batched_input_value(
|
||||
source_request: GenerationRequest,
|
||||
@@ -509,6 +617,7 @@ __all__ = [
|
||||
"load_generator_config_from_file",
|
||||
"normalize_generation_request",
|
||||
"normalize_generator_config",
|
||||
"register_continuation_kind",
|
||||
"request_to_pipeline_overrides",
|
||||
"request_to_sampling_param",
|
||||
]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
@@ -5,7 +5,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig, WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
@@ -11,10 +11,34 @@ import torch
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5ArchConfig(T5ArchConfig):
|
||||
"""T5 arch that pads tokenizer output to ``max_length``.
|
||||
|
||||
LongCat's denoising stage concatenates positive and negative
|
||||
attention masks along the batch dimension for CFG, which requires
|
||||
uniform seq length. The shared :class:`T5ArchConfig` dropped the
|
||||
``"padding": "max_length"`` tokenizer kwarg so other DiTs could run
|
||||
with variable-length masks; LongCat still needs the uniform
|
||||
contract.
|
||||
"""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5Config(T5Config):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LongCatT5ArchConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields."""
|
||||
@@ -103,8 +127,9 @@ class LongCatT2V480PConfig(PipelineConfig):
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
# UMT5 uses T5-like config; postprocess pads to 512. LongCatT5Config
|
||||
# restores ``padding="max_length"`` for the CFG concat contract.
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (LongCatT5Config(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (longcat_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (umt5_postprocess_text, ))
|
||||
|
||||
@@ -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
|
||||
@@ -1,4 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.entrypoints.streaming.server import run_server
|
||||
from fastvideo.entrypoints.streaming.server import build_app, run_server
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
BlobStore,
|
||||
InMemoryBlobStore,
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
|
||||
__all__ = ["run_server"]
|
||||
__all__ = [
|
||||
"BlobStore",
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionState",
|
||||
"SessionStore",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""JSON WebSocket protocol schemas for the streaming server.
|
||||
|
||||
Every control message shares the envelope ``{"type": <str>, ...}``.
|
||||
Pydantic models live here so the server can parse / validate incoming
|
||||
frames and emit well-typed outgoing frames without hand-rolled dicts.
|
||||
|
||||
The message catalogue matches the contract in
|
||||
``docs/design/server_contracts/streaming.md``; additions must land in
|
||||
both places in the same PR.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated, Any, Literal, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client → server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SessionInitV2(BaseModel):
|
||||
"""Opening frame the client sends after the WebSocket handshake."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
type: Literal["session_init_v2"]
|
||||
client_id: str | None = None
|
||||
preset: str | None = None
|
||||
preset_label: str | None = None
|
||||
curated_prompts: list[str] = Field(default_factory=list)
|
||||
initial_image: dict[str, Any] | None = None
|
||||
enhancement_enabled: bool = False
|
||||
auto_extension_enabled: bool = False
|
||||
loop_generation_enabled: bool = False
|
||||
single_clip_mode: bool = False
|
||||
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
|
||||
continuation_state: dict[str, Any] | None = None
|
||||
"""Optional ``{kind, payload}`` dict; hydrated into
|
||||
:class:`fastvideo.api.ContinuationState` server-side."""
|
||||
|
||||
|
||||
class SegmentPromptSource(BaseModel):
|
||||
"""Request a new segment using a specific prompt."""
|
||||
|
||||
type: Literal["segment_prompt_source"]
|
||||
prompt: str
|
||||
negative_prompt: str | None = None
|
||||
source: Literal["curated", "enhanced", "user", "auto_extension"] = "user"
|
||||
seed: int | None = None
|
||||
num_inference_steps: int | None = None
|
||||
guidance_scale: float | None = None
|
||||
|
||||
|
||||
class SeedPromptsUpdated(BaseModel):
|
||||
type: Literal["seed_prompts_updated"]
|
||||
seed_prompts: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class EnhancementUpdated(BaseModel):
|
||||
type: Literal["enhancement_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutoExtensionUpdated(BaseModel):
|
||||
type: Literal["auto_extension_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class LoopGenerationUpdated(BaseModel):
|
||||
type: Literal["loop_generation_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class GenerationPausedUpdated(BaseModel):
|
||||
type: Literal["generation_paused_updated"]
|
||||
paused: bool
|
||||
|
||||
|
||||
class SnapshotState(BaseModel):
|
||||
"""Request the current ``ContinuationState`` for export."""
|
||||
|
||||
type: Literal["snapshot_state"]
|
||||
|
||||
|
||||
ClientMessage = Annotated[
|
||||
Union[ # noqa: UP007 - Annotated requires Union for discriminator
|
||||
SessionInitV2,
|
||||
SegmentPromptSource,
|
||||
SeedPromptsUpdated,
|
||||
EnhancementUpdated,
|
||||
AutoExtensionUpdated,
|
||||
LoopGenerationUpdated,
|
||||
GenerationPausedUpdated,
|
||||
SnapshotState,
|
||||
],
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Server → client
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class QueueStatus(BaseModel):
|
||||
type: Literal["queue_status"] = "queue_status"
|
||||
position: int
|
||||
queue_depth: int
|
||||
|
||||
|
||||
class GpuAssigned(BaseModel):
|
||||
type: Literal["gpu_assigned"] = "gpu_assigned"
|
||||
gpu_id: int
|
||||
session_timeout: int
|
||||
|
||||
|
||||
class Ltx2StreamStart(BaseModel):
|
||||
type: Literal["ltx2_stream_start"] = "ltx2_stream_start"
|
||||
preset: str | None = None
|
||||
width: int
|
||||
height: int
|
||||
fps: int
|
||||
num_frames: int
|
||||
|
||||
|
||||
class Ltx2SegmentStart(BaseModel):
|
||||
type: Literal["ltx2_segment_start"] = "ltx2_segment_start"
|
||||
segment_idx: int
|
||||
prompt: str
|
||||
total_steps: int
|
||||
|
||||
|
||||
class StepComplete(BaseModel):
|
||||
type: Literal["step_complete"] = "step_complete"
|
||||
segment_idx: int
|
||||
step: int
|
||||
total_steps: int
|
||||
stage: str = "denoise"
|
||||
|
||||
|
||||
class MediaInit(BaseModel):
|
||||
"""Descriptor for the fMP4 initialization segment that follows."""
|
||||
|
||||
type: Literal["media_init"] = "media_init"
|
||||
segment_idx: int
|
||||
mime: str = "video/mp4; codecs=\"avc1.64001f, mp4a.40.2\""
|
||||
stream_id: str
|
||||
mode: Literal["av_fmp4"] = "av_fmp4"
|
||||
|
||||
|
||||
class MediaSegmentComplete(BaseModel):
|
||||
type: Literal["media_segment_complete"] = "media_segment_complete"
|
||||
segment_idx: int
|
||||
stream_id: str
|
||||
chunks: int
|
||||
duration_ms: float | None = None
|
||||
pts_base_ms: float | None = None
|
||||
|
||||
|
||||
class Ltx2SegmentComplete(BaseModel):
|
||||
type: Literal["ltx2_segment_complete"] = "ltx2_segment_complete"
|
||||
segment_idx: int
|
||||
generation_time_ms: float
|
||||
e2e_latency_ms: float | None = None
|
||||
|
||||
|
||||
class Ltx2StreamComplete(BaseModel):
|
||||
type: Literal["ltx2_stream_complete"] = "ltx2_stream_complete"
|
||||
reason: Literal["segment_cap", "stop_requested", "error"] = "stop_requested"
|
||||
|
||||
|
||||
class SessionTimeout(BaseModel):
|
||||
type: Literal["session_timeout"] = "session_timeout"
|
||||
timeout_seconds: int
|
||||
|
||||
|
||||
class ContinuationStateSnapshot(BaseModel):
|
||||
type: Literal["continuation_state_snapshot"] = "continuation_state_snapshot"
|
||||
state: dict[str, Any]
|
||||
"""``{kind, payload}`` dict matching
|
||||
:class:`fastvideo.api.ContinuationState`."""
|
||||
|
||||
|
||||
class ErrorMessage(BaseModel):
|
||||
type: Literal["error"] = "error"
|
||||
code: Literal[
|
||||
"session_rejected",
|
||||
"invalid_message",
|
||||
"preset_mismatch",
|
||||
"gpu_unavailable",
|
||||
"worker_failed",
|
||||
"upstream_timeout",
|
||||
"internal_error",
|
||||
] = "internal_error"
|
||||
message: str
|
||||
retryable: bool = False
|
||||
|
||||
|
||||
ServerMessage = Union[ # noqa: UP007 - pydantic Union handling
|
||||
QueueStatus,
|
||||
GpuAssigned,
|
||||
Ltx2StreamStart,
|
||||
Ltx2SegmentStart,
|
||||
StepComplete,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2StreamComplete,
|
||||
SessionTimeout,
|
||||
ContinuationStateSnapshot,
|
||||
ErrorMessage,
|
||||
]
|
||||
|
||||
|
||||
def parse_client_message(raw: dict[str, Any]) -> ClientMessage:
|
||||
"""Parse an incoming WebSocket dict into a typed client message.
|
||||
|
||||
Unknown ``type`` values raise :class:`pydantic.ValidationError`; the
|
||||
server handler turns that into an ``error`` frame with
|
||||
``code="invalid_message"``.
|
||||
"""
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
return TypeAdapter(ClientMessage).validate_python(raw)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AutoExtensionUpdated",
|
||||
"ClientMessage",
|
||||
"ContinuationStateSnapshot",
|
||||
"EnhancementUpdated",
|
||||
"ErrorMessage",
|
||||
"GenerationPausedUpdated",
|
||||
"GpuAssigned",
|
||||
"Ltx2SegmentComplete",
|
||||
"Ltx2SegmentStart",
|
||||
"Ltx2StreamComplete",
|
||||
"Ltx2StreamStart",
|
||||
"LoopGenerationUpdated",
|
||||
"MediaInit",
|
||||
"MediaSegmentComplete",
|
||||
"QueueStatus",
|
||||
"SeedPromptsUpdated",
|
||||
"SegmentPromptSource",
|
||||
"ServerMessage",
|
||||
"SessionInitV2",
|
||||
"SessionTimeout",
|
||||
"SnapshotState",
|
||||
"StepComplete",
|
||||
"parse_client_message",
|
||||
]
|
||||
@@ -1,15 +1,531 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-generator FastAPI + WebSocket streaming server."""
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.api.schema import ServeConfig
|
||||
import asyncio
|
||||
import contextlib
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
AutoExtensionUpdated,
|
||||
ContinuationStateSnapshot,
|
||||
EnhancementUpdated,
|
||||
ErrorMessage,
|
||||
GenerationPausedUpdated,
|
||||
GpuAssigned,
|
||||
LoopGenerationUpdated,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamComplete,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SeedPromptsUpdated,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image, )
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# RFC 6455 WebSocket close codes used by the server.
|
||||
_WS_CLOSE_UNSUPPORTED_DATA = 1003
|
||||
_WS_CLOSE_TRY_AGAIN_LATER = 1013
|
||||
|
||||
def run_server(serve_config: ServeConfig) -> None:
|
||||
"""Launch the streaming (WebSocket / Dynamo) server."""
|
||||
|
||||
class _GeneratorProto(Protocol):
|
||||
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
|
||||
|
||||
def generate(self, request: GenerationRequest) -> Any:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerState:
|
||||
serve_config: ServeConfig
|
||||
generator: _GeneratorProto
|
||||
sessions: SessionManager
|
||||
session_store: SessionStore
|
||||
|
||||
|
||||
def build_app(
|
||||
serve_config: ServeConfig,
|
||||
generator: _GeneratorProto,
|
||||
*,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI app used by :func:`run_server`.
|
||||
|
||||
Exposed so tests can drive the WebSocket endpoint in-process via
|
||||
``starlette.testclient.TestClient(app).websocket_connect(...)``.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
|
||||
"server; got None. Add a `streaming:` block to your serve config.")
|
||||
|
||||
sessions = SessionManager(
|
||||
segment_cap=serve_config.streaming.generation_segment_cap,
|
||||
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
|
||||
)
|
||||
state = ServerState(
|
||||
serve_config=serve_config,
|
||||
generator=generator,
|
||||
sessions=sessions,
|
||||
session_store=session_store or InMemorySessionStore(),
|
||||
)
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming")
|
||||
|
||||
@app.get("/health")
|
||||
async def _health() -> JSONResponse:
|
||||
return JSONResponse({
|
||||
"status": "ok",
|
||||
"sessions": len(state.sessions),
|
||||
"stream_mode": state.serve_config.streaming.stream_mode,
|
||||
})
|
||||
|
||||
@app.websocket("/v1/stream")
|
||||
async def _stream(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
try:
|
||||
session = state.sessions.create()
|
||||
except SessionRejected as exc:
|
||||
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
|
||||
return
|
||||
|
||||
try:
|
||||
await _handle_session(websocket, session, state)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("session %s: client disconnected", session.id[:8])
|
||||
except Exception: # pragma: no cover - defensive catch-all
|
||||
logger.exception("session %s: unhandled error", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
finally:
|
||||
_cleanup_session(session, state)
|
||||
|
||||
app.state.server_state = state
|
||||
return app
|
||||
|
||||
|
||||
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
|
||||
"""Launch the streaming server.
|
||||
|
||||
Boots a :class:`fastvideo.VideoGenerator` from
|
||||
``serve_config.generator`` unless ``generator`` is provided, then
|
||||
serves ``build_app(...)`` via uvicorn.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
|
||||
"got None. Add a `streaming:` block to your serve config.")
|
||||
raise NotImplementedError("streaming server is not implemented yet")
|
||||
|
||||
import uvicorn
|
||||
|
||||
if generator is None:
|
||||
from fastvideo import VideoGenerator # lazy to avoid boot cost
|
||||
|
||||
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
|
||||
app = build_app(serve_config, generator)
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
)
|
||||
|
||||
|
||||
async def _handle_session(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
init = await _read_init_message(websocket, session, state)
|
||||
if init is None:
|
||||
return
|
||||
|
||||
await _apply_session_init(session, init, state)
|
||||
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
|
||||
session.transition(SessionState.GPU_BINDING)
|
||||
await _send_json(websocket, GpuAssigned(
|
||||
gpu_id=0,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
session.transition(SessionState.ACTIVE)
|
||||
await _send_json(websocket, _build_stream_start(session, state))
|
||||
|
||||
try:
|
||||
await _run_segment_loop(websocket, session, state)
|
||||
finally:
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
|
||||
|
||||
|
||||
async def _read_init_message(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> SessionInitV2 | None:
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: init timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return None
|
||||
except WebSocketDisconnect:
|
||||
return None
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
|
||||
return None
|
||||
if not isinstance(parsed, SessionInitV2):
|
||||
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
|
||||
return None
|
||||
return parsed
|
||||
|
||||
|
||||
async def _reject_init(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
message: str,
|
||||
close_reason: str,
|
||||
) -> None:
|
||||
await _send_error(websocket, "invalid_message", message, retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.REJECTED)
|
||||
|
||||
|
||||
async def _apply_session_init(
|
||||
session: Session,
|
||||
init: SessionInitV2,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
session.client_id = init.client_id
|
||||
session.preset = init.preset
|
||||
session.preset_label = init.preset_label
|
||||
session.curated_prompts = list(init.curated_prompts)
|
||||
session.enhancement_enabled = init.enhancement_enabled
|
||||
session.auto_extension_enabled = init.auto_extension_enabled
|
||||
session.loop_generation_enabled = init.loop_generation_enabled
|
||||
session.single_clip_mode = init.single_clip_mode
|
||||
session.stream_mode = init.stream_mode
|
||||
|
||||
if init.initial_image is not None:
|
||||
# Decode + disk write off the event loop; payload is up to 32 MiB.
|
||||
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
|
||||
if image is not None:
|
||||
session.metadata["session_init_image"] = image.path
|
||||
|
||||
if init.continuation_state is not None:
|
||||
session.continuation_state = _coerce_state(init.continuation_state)
|
||||
if session.continuation_state is not None:
|
||||
state.session_store.store(session.id, session.continuation_state)
|
||||
|
||||
|
||||
async def _run_segment_loop(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
cap = state.sessions.segment_cap
|
||||
while True:
|
||||
if session.segment_cap_reached(cap):
|
||||
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
|
||||
return
|
||||
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: idle timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return
|
||||
except WebSocketDisconnect:
|
||||
return
|
||||
session.touch()
|
||||
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SnapshotState):
|
||||
snap = state.session_store.snapshot(session.id)
|
||||
if snap is None:
|
||||
await _send_error(websocket,
|
||||
"internal_error",
|
||||
"no continuation state available for session",
|
||||
retryable=False)
|
||||
continue
|
||||
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SegmentPromptSource):
|
||||
await _run_segment(websocket, session, state, parsed)
|
||||
continue
|
||||
|
||||
# Silently ignore unknown-but-valid types (additive-evolution
|
||||
# rule in streaming.md).
|
||||
_apply_toggle(session, parsed)
|
||||
|
||||
|
||||
async def _run_segment(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
message: SegmentPromptSource,
|
||||
) -> None:
|
||||
request = _build_generation_request(session, message, state)
|
||||
segment_idx = session.segment_idx
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentStart(
|
||||
segment_idx=segment_idx,
|
||||
prompt=message.prompt,
|
||||
total_steps=request.sampling.num_inference_steps,
|
||||
))
|
||||
|
||||
start = time.perf_counter()
|
||||
loop = asyncio.get_running_loop()
|
||||
# TODO: executor-wrapped generate() cannot be cancelled, so a
|
||||
# client disconnect mid-segment leaves the GPU work running to
|
||||
# completion. Real cancellation needs the generate_async API.
|
||||
try:
|
||||
result = await loop.run_in_executor(None, state.generator.generate, request)
|
||||
except Exception as exc:
|
||||
logger.exception("session %s: generator failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
elapsed_ms = (time.perf_counter() - start) * 1000.0
|
||||
|
||||
frames = _extract_frames(result)
|
||||
if not frames:
|
||||
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
|
||||
# Synchronous generator call has no per-step hook; emit one
|
||||
# terminal StepComplete so observability wiring still sees the
|
||||
# segment finish.
|
||||
total = request.sampling.num_inference_steps
|
||||
await _send_json(websocket, StepComplete(
|
||||
segment_idx=segment_idx,
|
||||
step=total,
|
||||
total_steps=total,
|
||||
stage="denoise",
|
||||
))
|
||||
|
||||
encoder = FragmentedMP4Encoder(
|
||||
width=request.sampling.width,
|
||||
height=request.sampling.height,
|
||||
fps=request.sampling.fps,
|
||||
segment_idx=segment_idx,
|
||||
)
|
||||
chunks_relayed = 0
|
||||
async with encoder:
|
||||
init_sent = False
|
||||
async for chunk in encoder.encode(frames):
|
||||
if chunk.kind == "init":
|
||||
await _send_json(websocket, MediaInit(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=chunk.stream_id,
|
||||
))
|
||||
init_sent = True
|
||||
await websocket.send_bytes(chunk.data)
|
||||
if init_sent and chunk.kind == "media":
|
||||
chunks_relayed += 1
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
MediaSegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=encoder.stream_id,
|
||||
chunks=chunks_relayed,
|
||||
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
|
||||
))
|
||||
|
||||
new_state = _extract_state(result)
|
||||
if new_state is not None:
|
||||
session.continuation_state = new_state
|
||||
state.session_store.store(session.id, new_state)
|
||||
|
||||
session.segment_idx += 1
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ACTIVE)
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
generation_time_ms=elapsed_ms,
|
||||
e2e_latency_ms=elapsed_ms,
|
||||
))
|
||||
|
||||
|
||||
def _build_stream_start(
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> Ltx2StreamStart:
|
||||
default = state.serve_config.default_request
|
||||
return Ltx2StreamStart(
|
||||
preset=session.preset,
|
||||
width=default.sampling.width,
|
||||
height=default.sampling.height,
|
||||
fps=default.sampling.fps,
|
||||
num_frames=default.sampling.num_frames,
|
||||
)
|
||||
|
||||
|
||||
def _build_generation_request(
|
||||
session: Session,
|
||||
message: SegmentPromptSource,
|
||||
state: ServerState,
|
||||
) -> GenerationRequest:
|
||||
# Start from the operator-pinned default_request to pick up the
|
||||
# preset-selected sampling knobs; override with per-message values.
|
||||
base = state.serve_config.default_request
|
||||
sampling_kwargs: dict[str, Any] = {
|
||||
"num_videos_per_prompt":
|
||||
base.sampling.num_videos_per_prompt,
|
||||
"seed":
|
||||
message.seed if message.seed is not None else base.sampling.seed,
|
||||
"num_frames":
|
||||
base.sampling.num_frames,
|
||||
"height":
|
||||
base.sampling.height,
|
||||
"width":
|
||||
base.sampling.width,
|
||||
"fps":
|
||||
base.sampling.fps,
|
||||
"num_inference_steps":
|
||||
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
|
||||
"guidance_scale":
|
||||
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
|
||||
}
|
||||
request = GenerationRequest(
|
||||
prompt=message.prompt,
|
||||
negative_prompt=message.negative_prompt or base.negative_prompt,
|
||||
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
|
||||
sampling=SamplingConfig(**sampling_kwargs),
|
||||
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
|
||||
state=session.continuation_state,
|
||||
)
|
||||
return request
|
||||
|
||||
|
||||
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
|
||||
kind = raw.get("kind")
|
||||
payload = raw.get("payload")
|
||||
if not isinstance(kind, str) or not isinstance(payload, dict):
|
||||
return None
|
||||
return ContinuationState(kind=kind, payload=payload)
|
||||
|
||||
|
||||
def _apply_toggle(session: Session, message: Any) -> None:
|
||||
if isinstance(message, EnhancementUpdated):
|
||||
session.enhancement_enabled = message.enabled
|
||||
elif isinstance(message, AutoExtensionUpdated):
|
||||
session.auto_extension_enabled = message.enabled
|
||||
elif isinstance(message, LoopGenerationUpdated):
|
||||
session.loop_generation_enabled = message.enabled
|
||||
elif isinstance(message, GenerationPausedUpdated):
|
||||
session.generation_paused = message.paused
|
||||
elif isinstance(message, SeedPromptsUpdated):
|
||||
session.curated_prompts = list(message.seed_prompts)
|
||||
|
||||
|
||||
def _extract_frames(result: Any) -> list:
|
||||
if hasattr(result, "frames"):
|
||||
return list(result.frames or [])
|
||||
if isinstance(result, dict):
|
||||
return list(result.get("frames") or [])
|
||||
return []
|
||||
|
||||
|
||||
def _extract_state(result: Any) -> ContinuationState | None:
|
||||
state = getattr(result, "state", None)
|
||||
if state is None and isinstance(result, dict):
|
||||
state = result.get("state")
|
||||
if isinstance(state, ContinuationState):
|
||||
return state
|
||||
if isinstance(state, dict):
|
||||
return _coerce_state(state)
|
||||
return None
|
||||
|
||||
|
||||
async def _send_json(websocket: WebSocket, message: Any) -> None:
|
||||
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
|
||||
await websocket.send_json(payload)
|
||||
|
||||
|
||||
async def _send_error(
|
||||
websocket: WebSocket,
|
||||
code: str,
|
||||
message: str,
|
||||
*,
|
||||
retryable: bool,
|
||||
) -> None:
|
||||
await _send_json(
|
||||
websocket,
|
||||
ErrorMessage(code=code, message=message, retryable=retryable),
|
||||
)
|
||||
|
||||
|
||||
def _cleanup_session(session: Session, state: ServerState) -> None:
|
||||
state.sessions.close(session.id)
|
||||
state.session_store.drop(session.id)
|
||||
init_image_path = session.metadata.get("session_init_image")
|
||||
if isinstance(init_image_path, str):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(init_image_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ServerState",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-connection session lifecycle for the streaming server.
|
||||
|
||||
Each WebSocket opens exactly one :class:`Session`. :class:`SessionManager`
|
||||
enforces the ``generation_segment_cap`` and ``session_timeout_seconds``
|
||||
budgets from :class:`fastvideo.api.StreamingConfig`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
class SessionState(enum.Enum):
|
||||
"""State-machine positions for a streaming session.
|
||||
|
||||
Transitions are server-owned. See
|
||||
``docs/design/server_contracts/streaming.md`` for the full diagram.
|
||||
"""
|
||||
|
||||
INITIALIZING = "initializing"
|
||||
QUEUED = "queued"
|
||||
GPU_BINDING = "gpu_binding"
|
||||
ACTIVE = "active"
|
||||
COMPLETE = "complete"
|
||||
ERROR = "error"
|
||||
TIMEOUT = "timeout"
|
||||
REJECTED = "rejected"
|
||||
|
||||
|
||||
_VALID_TRANSITIONS: dict[SessionState, frozenset[SessionState]] = {
|
||||
SessionState.INITIALIZING:
|
||||
frozenset({
|
||||
SessionState.QUEUED,
|
||||
SessionState.GPU_BINDING,
|
||||
SessionState.REJECTED,
|
||||
SessionState.ERROR,
|
||||
}),
|
||||
SessionState.QUEUED:
|
||||
frozenset({
|
||||
SessionState.GPU_BINDING,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
SessionState.REJECTED,
|
||||
}),
|
||||
SessionState.GPU_BINDING:
|
||||
frozenset({
|
||||
SessionState.ACTIVE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
}),
|
||||
SessionState.ACTIVE:
|
||||
frozenset({
|
||||
SessionState.ACTIVE,
|
||||
SessionState.COMPLETE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
}),
|
||||
SessionState.COMPLETE:
|
||||
frozenset(),
|
||||
SessionState.ERROR:
|
||||
frozenset(),
|
||||
SessionState.TIMEOUT:
|
||||
frozenset(),
|
||||
SessionState.REJECTED:
|
||||
frozenset(),
|
||||
}
|
||||
|
||||
|
||||
class InvalidSessionTransition(RuntimeError):
|
||||
"""Raised when a session is asked to transition along an illegal edge."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Session:
|
||||
id: str = field(default_factory=lambda: uuid.uuid4().hex)
|
||||
state: SessionState = SessionState.INITIALIZING
|
||||
created_at: float = field(default_factory=time.monotonic)
|
||||
last_activity: float = field(default_factory=time.monotonic)
|
||||
|
||||
client_id: str | None = None
|
||||
preset: str | None = None
|
||||
preset_label: str | None = None
|
||||
|
||||
curated_prompts: list[str] = field(default_factory=list)
|
||||
|
||||
segment_idx: int = 0
|
||||
|
||||
enhancement_enabled: bool = False
|
||||
auto_extension_enabled: bool = False
|
||||
loop_generation_enabled: bool = False
|
||||
single_clip_mode: bool = False
|
||||
generation_paused: bool = False
|
||||
|
||||
stream_mode: str = "av_fmp4"
|
||||
gpu_id: int | None = None
|
||||
|
||||
continuation_state: ContinuationState | None = None
|
||||
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def transition(self, target: SessionState) -> None:
|
||||
"""Move to ``target`` if the edge is allowed.
|
||||
|
||||
Raises :class:`InvalidSessionTransition` on illegal moves. The
|
||||
self-loop on ``ACTIVE`` is legal so the server can re-assert
|
||||
ACTIVE on segment completion without special casing.
|
||||
"""
|
||||
allowed = _VALID_TRANSITIONS.get(self.state, frozenset())
|
||||
if target not in allowed and target is not self.state:
|
||||
raise InvalidSessionTransition(f"{self.state.value} -> {target.value} is not a valid "
|
||||
f"session transition")
|
||||
self.state = target
|
||||
self.last_activity = time.monotonic()
|
||||
|
||||
def touch(self) -> None:
|
||||
self.last_activity = time.monotonic()
|
||||
|
||||
def is_active(self) -> bool:
|
||||
return self.state is SessionState.ACTIVE
|
||||
|
||||
def segment_cap_reached(self, cap: int) -> bool:
|
||||
return self.segment_idx >= cap
|
||||
|
||||
|
||||
class SessionManager:
|
||||
"""Registers sessions and enforces per-server session limits."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
segment_cap: int,
|
||||
session_timeout_seconds: int,
|
||||
max_sessions: int = 1,
|
||||
) -> None:
|
||||
self._segment_cap = segment_cap
|
||||
self._session_timeout_seconds = session_timeout_seconds
|
||||
self._max_sessions = max_sessions
|
||||
self._sessions: dict[str, Session] = {}
|
||||
|
||||
@property
|
||||
def segment_cap(self) -> int:
|
||||
return self._segment_cap
|
||||
|
||||
@property
|
||||
def session_timeout_seconds(self) -> int:
|
||||
return self._session_timeout_seconds
|
||||
|
||||
def create(self) -> Session:
|
||||
if len(self._sessions) >= self._max_sessions:
|
||||
raise SessionRejected(f"max sessions reached ({self._max_sessions})")
|
||||
session = Session()
|
||||
self._sessions[session.id] = session
|
||||
return session
|
||||
|
||||
def get(self, session_id: str) -> Session | None:
|
||||
return self._sessions.get(session_id)
|
||||
|
||||
def close(self, session_id: str) -> None:
|
||||
self._sessions.pop(session_id, None)
|
||||
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
return session_id in self._sessions
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._sessions)
|
||||
|
||||
def active_sessions(self) -> list[Session]:
|
||||
return [s for s in self._sessions.values() if s.is_active()]
|
||||
|
||||
def reap_timed_out(self, now: float | None = None) -> list[str]:
|
||||
"""Return the ids of sessions that have exceeded the idle timeout.
|
||||
|
||||
The caller is responsible for actually closing them — this
|
||||
method only *identifies* dead sessions so the server can emit
|
||||
``session_timeout`` frames before dropping the WebSocket.
|
||||
|
||||
TODO: unused until a background driver calls it. Per-connection
|
||||
idle enforcement currently happens via asyncio.wait_for on
|
||||
receive_json; this helper catches sessions stuck before any
|
||||
receive (e.g. future QUEUED state) and is expected to be wired
|
||||
into the GPU-pool reaper.
|
||||
"""
|
||||
now = now if now is not None else time.monotonic()
|
||||
dead: list[str] = []
|
||||
for sid, session in self._sessions.items():
|
||||
if session.state in {
|
||||
SessionState.COMPLETE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
SessionState.REJECTED,
|
||||
}:
|
||||
continue
|
||||
if now - session.last_activity > self._session_timeout_seconds:
|
||||
dead.append(sid)
|
||||
return dead
|
||||
|
||||
|
||||
class SessionRejected(RuntimeError):
|
||||
"""Raised when session creation fails (queue full, auth, etc.)."""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InvalidSessionTransition",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionRejected",
|
||||
"SessionState",
|
||||
]
|
||||
@@ -0,0 +1,103 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Persist the initial-image blob attached to a streaming session."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import contextlib
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
_ACCEPTED_MIMES = {
|
||||
"image/png": ".png",
|
||||
"image/jpeg": ".jpg",
|
||||
"image/jpg": ".jpg",
|
||||
"image/webp": ".webp",
|
||||
}
|
||||
|
||||
_MAX_IMAGE_BYTES = 32 * 1024 * 1024 # 32 MiB cap
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SessionInitImage:
|
||||
"""Location of the persisted init image.
|
||||
|
||||
Callers pass ``path`` to ``InputConfig.image_path``; ``display_name``
|
||||
is only used for logs.
|
||||
"""
|
||||
|
||||
path: str
|
||||
display_name: str
|
||||
mime: str
|
||||
|
||||
|
||||
def persist_session_init_image(
|
||||
payload: Any,
|
||||
*,
|
||||
output_dir: str | None = None,
|
||||
) -> SessionInitImage | None:
|
||||
"""Decode a client init-image blob and persist it to disk.
|
||||
|
||||
``payload`` shape (matches the internal UI protocol)::
|
||||
|
||||
{
|
||||
"mime": "image/png",
|
||||
"name": "ref.png",
|
||||
"data": "<base64 bytes>",
|
||||
}
|
||||
|
||||
Returns ``None`` when ``payload`` is falsy (no init image). Raises
|
||||
:class:`ValueError` on schema / size / decode errors so the caller
|
||||
can surface a user-facing ``error`` frame.
|
||||
"""
|
||||
if not payload:
|
||||
return None
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("session init image must be an object")
|
||||
|
||||
mime = payload.get("mime")
|
||||
if mime not in _ACCEPTED_MIMES:
|
||||
raise ValueError(f"session init image mime {mime!r} is not one of "
|
||||
f"{sorted(_ACCEPTED_MIMES)}")
|
||||
data_b64 = payload.get("data")
|
||||
if not isinstance(data_b64, str):
|
||||
raise ValueError("session init image data must be a base64 string")
|
||||
try:
|
||||
data = base64.b64decode(data_b64, validate=True)
|
||||
except (binascii.Error, ValueError) as exc:
|
||||
raise ValueError(f"session init image data is not valid base64: {exc}") from exc
|
||||
if len(data) > _MAX_IMAGE_BYTES:
|
||||
raise ValueError(f"session init image is {len(data)} bytes; limit is "
|
||||
f"{_MAX_IMAGE_BYTES}")
|
||||
if len(data) == 0:
|
||||
raise ValueError("session init image data is empty")
|
||||
|
||||
ext = _ACCEPTED_MIMES[mime]
|
||||
display_name = _sanitize_display_name(payload.get("name")) or f"init{ext}"
|
||||
fd, path = tempfile.mkstemp(prefix="fastvideo-init-", suffix=ext, dir=output_dir)
|
||||
try:
|
||||
with os.fdopen(fd, "wb") as f:
|
||||
f.write(data)
|
||||
except Exception:
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(path)
|
||||
raise
|
||||
return SessionInitImage(path=path, display_name=display_name, mime=mime)
|
||||
|
||||
|
||||
def _sanitize_display_name(name: Any) -> str | None:
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
name = name.strip()
|
||||
if not name:
|
||||
return None
|
||||
# Strip any path components — we only keep the leaf for logging.
|
||||
return os.path.basename(name)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"SessionInitImage",
|
||||
"persist_session_init_image",
|
||||
]
|
||||
@@ -0,0 +1,206 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Session state store for the FastVideo streaming server.
|
||||
|
||||
The streaming server keeps continuation state (decoded frames + audio
|
||||
latents from the previous segment) server-side so the client doesn't
|
||||
re-upload multi-megabyte tensors each WebSocket message. Two operations
|
||||
are needed:
|
||||
|
||||
* ``snapshot(session_id) -> ContinuationState`` — serialize the current
|
||||
state so it can be exported (e.g. over HTTP) or migrated to a
|
||||
different server.
|
||||
* ``hydrate(state) -> session_id`` — load a previously serialized state
|
||||
into a new session (for resume-after-disconnect flows).
|
||||
|
||||
The store is an ABC with an :class:`InMemorySessionStore` default; Redis
|
||||
or other backends can drop in without touching the pipeline.
|
||||
|
||||
Large tensor payloads (video frames, audio latents) are kept out of the
|
||||
JSON payload via an accompanying :class:`BlobStore`. Both stores share a
|
||||
process today; they are separate types so that a future implementation
|
||||
can put blobs on S3 while keeping session metadata in Redis.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
class BlobStore(ABC):
|
||||
"""Opaque byte-blob storage keyed by id.
|
||||
|
||||
A :class:`ContinuationState` payload can reference large tensors
|
||||
stored in a :class:`BlobStore` rather than inlining them, so the
|
||||
JSON payload stays small when the state travels over the wire.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
|
||||
"""Store ``data`` and return a blob id for later retrieval."""
|
||||
|
||||
@abstractmethod
|
||||
def get(self, blob_id: str) -> bytes:
|
||||
"""Load a previously stored blob. Raises ``KeyError`` if absent."""
|
||||
|
||||
@abstractmethod
|
||||
def drop(self, blob_id: str) -> None:
|
||||
"""Remove a blob. Missing ids are a no-op."""
|
||||
|
||||
@abstractmethod
|
||||
def __contains__(self, blob_id: str) -> bool:
|
||||
...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _BlobRecord:
|
||||
data: bytes
|
||||
mime: str
|
||||
|
||||
|
||||
class InMemoryBlobStore(BlobStore):
|
||||
"""Thread-safe in-memory :class:`BlobStore` for single-process servers.
|
||||
|
||||
No eviction policy — callers are responsible for calling
|
||||
:meth:`drop` when a blob's owning state is replaced or a session
|
||||
ends. A redis- or filesystem-backed :class:`BlobStore` should
|
||||
replace this when the streaming server lands as a real service
|
||||
(PR 7.5+).
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._blobs: dict[str, _BlobRecord] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
|
||||
blob_id = uuid.uuid4().hex
|
||||
with self._lock:
|
||||
self._blobs[blob_id] = _BlobRecord(data=data, mime=mime)
|
||||
return blob_id
|
||||
|
||||
def get(self, blob_id: str) -> bytes:
|
||||
with self._lock:
|
||||
record = self._blobs.get(blob_id)
|
||||
if record is None:
|
||||
raise KeyError(f"Unknown blob id: {blob_id}")
|
||||
return record.data
|
||||
|
||||
def drop(self, blob_id: str) -> None:
|
||||
with self._lock:
|
||||
self._blobs.pop(blob_id, None)
|
||||
|
||||
def __contains__(self, blob_id: str) -> bool:
|
||||
with self._lock:
|
||||
return blob_id in self._blobs
|
||||
|
||||
def __len__(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._blobs)
|
||||
|
||||
|
||||
class SessionStore(ABC):
|
||||
"""Keyed store for per-session continuation state.
|
||||
|
||||
Implementations own the session-id → state mapping. The streaming
|
||||
server calls :meth:`store` after each segment and :meth:`snapshot`
|
||||
when a client explicitly asks for an exportable state handle.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def store(self, session_id: str, state: ContinuationState) -> None:
|
||||
"""Persist ``state`` for ``session_id``, replacing any prior value."""
|
||||
|
||||
@abstractmethod
|
||||
def snapshot(self, session_id: str) -> ContinuationState | None:
|
||||
"""Return the current state for ``session_id`` (or ``None``)."""
|
||||
|
||||
@abstractmethod
|
||||
def hydrate(
|
||||
self,
|
||||
state: ContinuationState,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
) -> str:
|
||||
"""Install ``state`` as the starting point for a session.
|
||||
|
||||
When ``session_id`` is ``None`` the store allocates a fresh id
|
||||
(UUID4); when provided the store uses it verbatim, overwriting
|
||||
any prior state at that id.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def drop(self, session_id: str) -> None:
|
||||
"""Forget a session. Missing ids are a no-op."""
|
||||
|
||||
@abstractmethod
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def __iter__(self) -> Iterator[str]:
|
||||
...
|
||||
|
||||
|
||||
class InMemorySessionStore(SessionStore):
|
||||
"""Thread-safe in-memory :class:`SessionStore`.
|
||||
|
||||
Default implementation used by single-process deployments; a future
|
||||
Redis-backed store can be dropped in without changes to the server.
|
||||
|
||||
No eviction / TTL / bounded capacity — sessions only leave via
|
||||
:meth:`drop`. The live streaming server (PR 7.5+) is responsible
|
||||
for bounding growth and for dropping any :class:`BlobStore` blobs
|
||||
referenced by a state when that state is replaced or a session
|
||||
ends; this class does not know about blobs.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._sessions: dict[str, ContinuationState] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def store(self, session_id: str, state: ContinuationState) -> None:
|
||||
with self._lock:
|
||||
self._sessions[session_id] = state
|
||||
|
||||
def snapshot(self, session_id: str) -> ContinuationState | None:
|
||||
with self._lock:
|
||||
return self._sessions.get(session_id)
|
||||
|
||||
def hydrate(
|
||||
self,
|
||||
state: ContinuationState,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
) -> str:
|
||||
sid = session_id or uuid.uuid4().hex
|
||||
with self._lock:
|
||||
self._sessions[sid] = state
|
||||
return sid
|
||||
|
||||
def drop(self, session_id: str) -> None:
|
||||
with self._lock:
|
||||
self._sessions.pop(session_id, None)
|
||||
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
with self._lock:
|
||||
return session_id in self._sessions
|
||||
|
||||
def __iter__(self) -> Iterator[str]:
|
||||
with self._lock:
|
||||
return iter(list(self._sessions))
|
||||
|
||||
def __len__(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._sessions)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BlobStore",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"SessionStore",
|
||||
]
|
||||
@@ -0,0 +1,213 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""fMP4 stream encoder used by the streaming server.
|
||||
|
||||
The client's Media Source Extensions player needs a continuous fMP4
|
||||
byte stream: first an *initialization segment* (``ftyp`` + ``moov``),
|
||||
then one or more *media segments* (``moof`` + ``mdat``). We pipe raw
|
||||
RGB frames into an ffmpeg subprocess configured for fragmented output
|
||||
via ``-movflags empty_moov+default_base_moof+frag_keyframe+faststart``
|
||||
and stream the bytes back out.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import subprocess
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import numpy as np
|
||||
|
||||
|
||||
@dataclass
|
||||
class FragmentedMP4Chunk:
|
||||
"""A single fMP4 byte chunk emitted by :class:`FragmentedMP4Encoder`.
|
||||
|
||||
``kind`` identifies whether the chunk is the init segment (must be
|
||||
fed into the client's ``SourceBuffer`` first) or a media fragment.
|
||||
"""
|
||||
|
||||
kind: Literal["init", "media"]
|
||||
data: bytes
|
||||
stream_id: str
|
||||
segment_idx: int
|
||||
|
||||
|
||||
class FragmentedMP4Encoder:
|
||||
"""Stream RGB frames in, fMP4 chunks out.
|
||||
|
||||
One encoder covers one segment. The server creates a new encoder
|
||||
per :class:`ltx2_segment_start`` boundary so each segment becomes
|
||||
one media fragment the client can append independently.
|
||||
|
||||
Example::
|
||||
|
||||
encoder = FragmentedMP4Encoder(width=1024, height=576, fps=24,
|
||||
segment_idx=0)
|
||||
async with encoder:
|
||||
async for chunk in encoder.encode(frames):
|
||||
await websocket.send_bytes(chunk.data)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
width: int,
|
||||
height: int,
|
||||
fps: int,
|
||||
segment_idx: int,
|
||||
stream_id: str | None = None,
|
||||
ffmpeg_path: str = "ffmpeg",
|
||||
preset: str = "ultrafast",
|
||||
pixel_format_out: str = "yuv420p",
|
||||
extra_args: list[str] | None = None,
|
||||
) -> None:
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.fps = fps
|
||||
self.segment_idx = segment_idx
|
||||
self.stream_id = stream_id or uuid.uuid4().hex
|
||||
self._ffmpeg_path = ffmpeg_path
|
||||
self._preset = preset
|
||||
self._pixel_format_out = pixel_format_out
|
||||
self._extra_args = list(extra_args or [])
|
||||
self._proc: subprocess.Popen | None = None
|
||||
self._init_emitted = False
|
||||
|
||||
async def __aenter__(self) -> FragmentedMP4Encoder:
|
||||
self._spawn()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None:
|
||||
await self.close()
|
||||
|
||||
def _spawn(self) -> None:
|
||||
args = [
|
||||
self._ffmpeg_path,
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"error",
|
||||
"-f",
|
||||
"rawvideo",
|
||||
"-pix_fmt",
|
||||
"rgb24",
|
||||
"-s",
|
||||
f"{self.width}x{self.height}",
|
||||
"-r",
|
||||
str(self.fps),
|
||||
"-i",
|
||||
"-",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
self._preset,
|
||||
"-tune",
|
||||
"zerolatency",
|
||||
"-pix_fmt",
|
||||
self._pixel_format_out,
|
||||
"-movflags",
|
||||
"empty_moov+default_base_moof+frag_keyframe+faststart",
|
||||
"-f",
|
||||
"mp4",
|
||||
*self._extra_args,
|
||||
"-",
|
||||
]
|
||||
# stderr → DEVNULL: with -loglevel error on, the only thing
|
||||
# stderr would carry is unsolicited warnings. Piping without a
|
||||
# reader deadlocks ffmpeg once the pipe buffer (~64 KiB) fills.
|
||||
self._proc = subprocess.Popen( # noqa: S603
|
||||
args,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
bufsize=0,
|
||||
)
|
||||
|
||||
async def encode(
|
||||
self,
|
||||
frames: list[np.ndarray] | AsyncIterator[np.ndarray],
|
||||
) -> AsyncIterator[FragmentedMP4Chunk]:
|
||||
"""Feed frames into ffmpeg and yield fMP4 chunks as they appear."""
|
||||
if self._proc is None:
|
||||
self._spawn()
|
||||
assert self._proc is not None and self._proc.stdin is not None
|
||||
proc = self._proc
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
async def _writer() -> None:
|
||||
try:
|
||||
if hasattr(frames, "__aiter__"):
|
||||
async for frame in frames: # type: ignore[union-attr]
|
||||
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
|
||||
else:
|
||||
for frame in frames: # type: ignore[assignment]
|
||||
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
|
||||
finally:
|
||||
with contextlib.suppress(BrokenPipeError):
|
||||
proc.stdin.close()
|
||||
|
||||
writer_task = asyncio.create_task(_writer())
|
||||
try:
|
||||
reader = proc.stdout
|
||||
assert reader is not None
|
||||
# Read in reasonably-sized chunks; MSE tolerates any size
|
||||
# but we don't want to starve the event loop.
|
||||
chunk_size = 64 * 1024
|
||||
while True:
|
||||
data = await loop.run_in_executor(None, reader.read, chunk_size)
|
||||
if not data:
|
||||
break
|
||||
kind: Literal["init", "media"] = "init" if not self._init_emitted else "media"
|
||||
self._init_emitted = True
|
||||
yield FragmentedMP4Chunk(
|
||||
kind=kind,
|
||||
data=bytes(data),
|
||||
stream_id=self.stream_id,
|
||||
segment_idx=self.segment_idx,
|
||||
)
|
||||
finally:
|
||||
await writer_task
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._proc is None:
|
||||
return
|
||||
proc = self._proc
|
||||
self._proc = None
|
||||
try:
|
||||
if proc.stdin and not proc.stdin.closed:
|
||||
proc.stdin.close()
|
||||
except BrokenPipeError:
|
||||
pass
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
loop.run_in_executor(None, proc.wait),
|
||||
timeout=5.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
proc.kill()
|
||||
await loop.run_in_executor(None, proc.wait)
|
||||
|
||||
|
||||
def _write_frame(stdin, frame: np.ndarray) -> None:
|
||||
import numpy as np
|
||||
|
||||
if not isinstance(frame, np.ndarray):
|
||||
raise TypeError("fMP4 encoder frames must be numpy.ndarray")
|
||||
if frame.dtype != np.uint8:
|
||||
frame = frame.astype(np.uint8)
|
||||
if frame.ndim != 3 or frame.shape[-1] != 3:
|
||||
raise ValueError("fMP4 encoder frames must be HxWx3 uint8 RGB; got "
|
||||
f"shape={frame.shape}, dtype={frame.dtype}")
|
||||
with contextlib.suppress(BrokenPipeError):
|
||||
stdin.write(frame.tobytes())
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -519,7 +519,7 @@ def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch)
|
||||
cli_main.main()
|
||||
|
||||
|
||||
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
|
||||
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, monkeypatch):
|
||||
config_path = tmp_path / "serve-streaming.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
@@ -530,9 +530,21 @@ def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
|
||||
)
|
||||
args, _ = _parse_serve_args(["--config", str(config_path)])
|
||||
|
||||
with pytest.raises(NotImplementedError,
|
||||
match="streaming server is not implemented"):
|
||||
ServeSubcommand().cmd(args)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_run_server(serve_config, *, generator=None):
|
||||
captured["serve_config"] = serve_config
|
||||
|
||||
def fail_if_called(*_args, **_kwargs):
|
||||
raise AssertionError("OpenAI server must not run when streaming is set")
|
||||
|
||||
monkeypatch.setattr(streaming_server, "run_server", fake_run_server)
|
||||
monkeypatch.setattr(api_server, "run_server", fail_if_called)
|
||||
ServeSubcommand().cmd(args)
|
||||
|
||||
serve_config = captured["serve_config"]
|
||||
assert serve_config.streaming is not None
|
||||
assert serve_config.streaming.stream_mode == "av_fmp4"
|
||||
|
||||
|
||||
def test_streaming_run_server_rejects_missing_streaming_block():
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for ``fastvideo.api.compat`` translation helpers covering the
|
||||
typed CompileConfig + PipelineSelection.vae_tiling surfaces promoted in
|
||||
PR 6.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.api.compat import (
|
||||
generator_config_to_fastvideo_args,
|
||||
legacy_from_pretrained_to_config,
|
||||
)
|
||||
from fastvideo.api.schema import CompileConfig, GeneratorConfig
|
||||
|
||||
|
||||
class TestLegacyTorchCompileKwargsTranslation:
|
||||
"""Legacy ``torch_compile_kwargs={...}`` gets split across the four
|
||||
first-class :class:`CompileConfig` fields and anything unknown falls
|
||||
into ``extras``."""
|
||||
|
||||
def test_all_typed_keys_promoted(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{
|
||||
"enable_torch_compile": True,
|
||||
"torch_compile_kwargs": {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "max-autotune-no-cudagraphs",
|
||||
"dynamic": False,
|
||||
},
|
||||
},
|
||||
)
|
||||
compile_config = config.engine.compile
|
||||
assert compile_config.enabled is True
|
||||
assert compile_config.backend == "inductor"
|
||||
assert compile_config.fullgraph is True
|
||||
assert compile_config.mode == "max-autotune-no-cudagraphs"
|
||||
assert compile_config.dynamic is False
|
||||
assert compile_config.extras == {}
|
||||
|
||||
def test_unknown_keys_land_in_extras(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{
|
||||
"enable_torch_compile": True,
|
||||
"torch_compile_kwargs": {
|
||||
"backend": "inductor",
|
||||
"options": {"triton.cudagraphs": False},
|
||||
"disable": False,
|
||||
},
|
||||
},
|
||||
)
|
||||
compile_config = config.engine.compile
|
||||
assert compile_config.backend == "inductor"
|
||||
assert compile_config.extras == {
|
||||
"options": {"triton.cudagraphs": False},
|
||||
"disable": False,
|
||||
}
|
||||
|
||||
def test_empty_kwargs_produces_empty_extras(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"torch_compile_kwargs": {}},
|
||||
)
|
||||
compile_config = config.engine.compile
|
||||
assert compile_config.extras == {}
|
||||
assert compile_config.backend is None
|
||||
|
||||
|
||||
class TestCompileConfigRoundTrip:
|
||||
"""typed CompileConfig -> FastVideoArgs.torch_compile_kwargs
|
||||
reconstruction drops ``None`` typed fields and merges ``extras``."""
|
||||
|
||||
def test_only_typed_fields_emitted(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(
|
||||
CompileConfig(enabled=True, backend="inductor", fullgraph=True)),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert args.kwargs["enable_torch_compile"] is True
|
||||
assert args.kwargs["torch_compile_kwargs"] == {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
}
|
||||
|
||||
def test_extras_merged_into_torch_compile_kwargs(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(
|
||||
CompileConfig(
|
||||
enabled=True,
|
||||
mode="reduce-overhead",
|
||||
extras={"options": {"triton.cudagraphs": False}},
|
||||
)),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert args.kwargs["torch_compile_kwargs"] == {
|
||||
"mode": "reduce-overhead",
|
||||
"options": {"triton.cudagraphs": False},
|
||||
}
|
||||
|
||||
def test_none_fields_suppressed(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(CompileConfig()),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert args.kwargs["torch_compile_kwargs"] == {}
|
||||
|
||||
|
||||
class TestLegacyLtx2VaeTilingTranslation:
|
||||
"""``ltx2_vae_tiling`` flat kwarg promotes to
|
||||
``generator.pipeline.vae_tiling``; reverse direction emits the
|
||||
legacy name back to FastVideoArgs."""
|
||||
|
||||
def test_forward_routes_to_pipeline_vae_tiling(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"ltx2_vae_tiling": False},
|
||||
)
|
||||
assert config.pipeline.vae_tiling is False
|
||||
|
||||
def test_true_round_trips(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"ltx2_vae_tiling": True},
|
||||
)
|
||||
assert config.pipeline.vae_tiling is True
|
||||
|
||||
def test_unset_stays_none(self) -> None:
|
||||
config = legacy_from_pretrained_to_config("/models/ltx2", {})
|
||||
assert config.pipeline.vae_tiling is None
|
||||
|
||||
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(CompileConfig()),
|
||||
)
|
||||
config.pipeline.vae_tiling = False
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert args.kwargs["ltx2_vae_tiling"] is False
|
||||
|
||||
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(CompileConfig()),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert "ltx2_vae_tiling" not in args.kwargs
|
||||
|
||||
|
||||
class TestLegacyTextEncoderCompileTranslation:
|
||||
"""``enable_torch_compile_text_encoder`` flat kwarg promotes to
|
||||
``generator.engine.compile.text_encoder_enabled``; reverse direction
|
||||
emits the legacy name back onto the FastVideoArgs kwargs dict so
|
||||
realtime-runtime consumers can read it before FastVideoArgs filters
|
||||
unknown fields."""
|
||||
|
||||
def test_forward_routes_to_compile_text_encoder_enabled(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"enable_torch_compile_text_encoder": True},
|
||||
)
|
||||
assert config.engine.compile.text_encoder_enabled is True
|
||||
|
||||
def test_false_round_trips(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"enable_torch_compile_text_encoder": False},
|
||||
)
|
||||
assert config.engine.compile.text_encoder_enabled is False
|
||||
|
||||
def test_unset_stays_none(self) -> None:
|
||||
config = legacy_from_pretrained_to_config("/models/ltx2", {})
|
||||
assert config.engine.compile.text_encoder_enabled is None
|
||||
|
||||
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(
|
||||
CompileConfig(text_encoder_enabled=True)),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert args.kwargs["enable_torch_compile_text_encoder"] is True
|
||||
|
||||
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(CompileConfig()),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert "enable_torch_compile_text_encoder" not in args.kwargs
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Helpers
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
def _engine_with_compile(compile_config):
|
||||
"""Build an ``EngineConfig`` that carries the supplied compile block."""
|
||||
from fastvideo.api.schema import EngineConfig
|
||||
engine = EngineConfig()
|
||||
engine.compile = compile_config
|
||||
return engine
|
||||
|
||||
|
||||
def _stub_fastvideo_args_from_kwargs(monkeypatch):
|
||||
"""Swap ``FastVideoArgs.from_kwargs`` for a capture-only stub so
|
||||
translation tests don't need to construct a valid FastVideoArgs."""
|
||||
from fastvideo import fastvideo_args as fva
|
||||
|
||||
class _Captured:
|
||||
|
||||
def __init__(self, **kw):
|
||||
self.kwargs = kw
|
||||
|
||||
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _Captured)
|
||||
@@ -0,0 +1,250 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the typed LTX-2 continuation state.
|
||||
|
||||
Covers:
|
||||
|
||||
* round-trip through :class:`ContinuationState` (inline and blob-backed)
|
||||
* payload is JSON-serializable (Dynamo RPC / HTTP client constraint)
|
||||
* kind / schema_version validation on deserialization
|
||||
* compat-layer validation (known kinds, payload shape)
|
||||
* round-trip through :func:`request_to_sampling_param` attaches the
|
||||
state to the resulting :class:`SamplingParam` without losing fidelity
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# Importing compat first, then the LTX-2 module, exercises the
|
||||
# self-registration side effect on import (important for the API
|
||||
# test suite where the pipeline package isn't otherwise imported).
|
||||
from fastvideo.api import compat as api_compat # noqa: F401
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
OutputConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import InMemoryBlobStore
|
||||
from fastvideo.pipelines.basic.ltx2.continuation import (
|
||||
LTX2_CONTINUATION_KIND,
|
||||
LTX2_CONTINUATION_SCHEMA_VERSION,
|
||||
LTX2ContinuationState,
|
||||
)
|
||||
|
||||
|
||||
def _make_typed_state() -> LTX2ContinuationState:
|
||||
return LTX2ContinuationState(
|
||||
segment_index=3,
|
||||
video_frames=[
|
||||
(np.ones((64, 64, 3), dtype=np.uint8) * (i * 10)) for i in range(4)
|
||||
],
|
||||
video_conditioning_frame_idx=9,
|
||||
video_conditioning_strength=0.75,
|
||||
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.float32),
|
||||
audio_sample_rate=24000,
|
||||
audio_conditioning_num_frames=5,
|
||||
audio_conditioning_strength=0.5,
|
||||
video_position_offset_sec=0.125,
|
||||
metadata={"note": "unit-test"},
|
||||
)
|
||||
|
||||
|
||||
class TestRoundTrip:
|
||||
"""Round-trip through :class:`ContinuationState` preserves all fields."""
|
||||
|
||||
def test_kind_and_schema_version(self):
|
||||
state = _make_typed_state().to_continuation_state()
|
||||
assert state.kind == LTX2_CONTINUATION_KIND
|
||||
assert state.payload["schema_version"] == LTX2_CONTINUATION_SCHEMA_VERSION
|
||||
|
||||
def test_inline_roundtrip_preserves_scalars(self):
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.segment_index == original.segment_index
|
||||
assert restored.video_conditioning_frame_idx == (
|
||||
original.video_conditioning_frame_idx)
|
||||
assert restored.video_conditioning_strength == (
|
||||
original.video_conditioning_strength)
|
||||
assert restored.audio_sample_rate == original.audio_sample_rate
|
||||
assert restored.audio_conditioning_num_frames == (
|
||||
original.audio_conditioning_num_frames)
|
||||
assert restored.audio_conditioning_strength == (
|
||||
original.audio_conditioning_strength)
|
||||
assert restored.video_position_offset_sec == (
|
||||
original.video_position_offset_sec)
|
||||
assert restored.metadata == original.metadata
|
||||
|
||||
def test_inline_roundtrip_preserves_video_frames(self):
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.video_frames is not None
|
||||
assert len(restored.video_frames) == len(original.video_frames)
|
||||
for before, after in zip(original.video_frames,
|
||||
restored.video_frames):
|
||||
np.testing.assert_array_equal(before, after)
|
||||
|
||||
def test_inline_roundtrip_preserves_audio_latents(self):
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.audio_latents is not None
|
||||
assert tuple(restored.audio_latents.shape) == tuple(
|
||||
original.audio_latents.shape)
|
||||
assert restored.audio_latents.dtype == original.audio_latents.dtype
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, original.audio_latents)
|
||||
|
||||
def test_payload_is_json_serializable(self):
|
||||
envelope = _make_typed_state().to_continuation_state()
|
||||
# json.dumps must not raise — required for Dynamo RPC transport
|
||||
# and HTTP client round-trip.
|
||||
reserialized = json.loads(json.dumps(envelope.payload))
|
||||
restored = LTX2ContinuationState.from_continuation_state(
|
||||
ContinuationState(
|
||||
kind=envelope.kind,
|
||||
payload=reserialized,
|
||||
))
|
||||
assert restored.segment_index == 3
|
||||
|
||||
def test_bf16_audio_latents_preserved(self):
|
||||
"""safetensors serialization must preserve bf16 dtype (numpy
|
||||
has no bf16, so a raw-bytes path would silently promote)."""
|
||||
state = LTX2ContinuationState(
|
||||
segment_index=0,
|
||||
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.bfloat16),
|
||||
)
|
||||
envelope = state.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.audio_latents is not None
|
||||
assert restored.audio_latents.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, state.audio_latents)
|
||||
|
||||
|
||||
class TestBlobIndirection:
|
||||
"""Large tensors live in the :class:`BlobStore` instead of the payload."""
|
||||
|
||||
def test_threshold_triggers_blob_path(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
state = _make_typed_state()
|
||||
envelope = state.to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=0,
|
||||
)
|
||||
assert "blob_id" in envelope.payload["video"]
|
||||
assert "blob_id" in envelope.payload["audio"]
|
||||
assert "frames_b64" not in envelope.payload["video"]
|
||||
assert "safetensors_b64" not in envelope.payload["audio"]
|
||||
assert len(blob_store) == 2
|
||||
|
||||
def test_blob_roundtrip_reconstructs_tensors(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=0,
|
||||
)
|
||||
restored = LTX2ContinuationState.from_continuation_state(
|
||||
envelope, blob_store=blob_store)
|
||||
assert restored.video_frames is not None
|
||||
assert len(restored.video_frames) == len(original.video_frames)
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, original.audio_latents)
|
||||
|
||||
def test_blob_id_held_when_store_unavailable(self):
|
||||
"""Deserializing without a blob store preserves the blob id so
|
||||
the caller can fetch it later."""
|
||||
blob_store = InMemoryBlobStore()
|
||||
envelope = _make_typed_state().to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=0,
|
||||
)
|
||||
blob_id_video = envelope.payload["video"]["blob_id"]
|
||||
blob_id_audio = envelope.payload["audio"]["blob_id"]
|
||||
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.video_frames is None
|
||||
assert restored.video_frames_blob_id == blob_id_video
|
||||
assert restored.audio_latents is None
|
||||
assert restored.audio_latents_blob_id == blob_id_audio
|
||||
|
||||
def test_large_threshold_keeps_payload_inline(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
envelope = _make_typed_state().to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=10 * 1024 * 1024, # 10 MiB
|
||||
)
|
||||
assert "frames_b64" in envelope.payload["video"]
|
||||
assert "safetensors_b64" in envelope.payload["audio"]
|
||||
assert len(blob_store) == 0
|
||||
|
||||
|
||||
class TestValidation:
|
||||
"""Invalid payloads error cleanly."""
|
||||
|
||||
def test_wrong_kind_rejected(self):
|
||||
envelope = ContinuationState(kind="longcat.v1", payload={})
|
||||
with pytest.raises(ValueError, match="Expected ContinuationState.kind"):
|
||||
LTX2ContinuationState.from_continuation_state(envelope)
|
||||
|
||||
def test_unsupported_schema_version_rejected(self):
|
||||
envelope = ContinuationState(
|
||||
kind=LTX2_CONTINUATION_KIND,
|
||||
payload={"schema_version": 999},
|
||||
)
|
||||
with pytest.raises(ValueError,
|
||||
match="Unsupported LTX-2 continuation schema"):
|
||||
LTX2ContinuationState.from_continuation_state(envelope)
|
||||
|
||||
def test_non_png_frame_rejected(self):
|
||||
state = LTX2ContinuationState(
|
||||
video_frames=[np.ones((64, 64, 3), dtype=np.float32)],
|
||||
)
|
||||
with pytest.raises(ValueError, match="uint8 HxWx3"):
|
||||
state.to_continuation_state()
|
||||
|
||||
|
||||
class TestCompatLayerWireUp:
|
||||
"""The public compat layer accepts request.state without reverting
|
||||
to NotImplementedError and attaches it to the SamplingParam path."""
|
||||
|
||||
def test_request_with_state_passes_through(self, tmp_path):
|
||||
# PR 7 removes the NotImplementedError for request.state; build a
|
||||
# minimal GenerationRequest carrying an LTX-2 state and make sure
|
||||
# the public boundary accepts it.
|
||||
from fastvideo.api.compat import (
|
||||
normalize_generation_request,
|
||||
_validate_continuation_state,
|
||||
)
|
||||
envelope = _make_typed_state().to_continuation_state()
|
||||
request = GenerationRequest(
|
||||
prompt="test",
|
||||
state=envelope,
|
||||
)
|
||||
normalized = normalize_generation_request(request)
|
||||
_validate_continuation_state(normalized.state)
|
||||
|
||||
def test_unknown_kind_rejected_at_boundary(self):
|
||||
from fastvideo.api.compat import _validate_continuation_state
|
||||
with pytest.raises(ValueError, match="Unknown ContinuationState kind"):
|
||||
_validate_continuation_state(
|
||||
ContinuationState(kind="mystery.v1", payload={}))
|
||||
|
||||
def test_empty_kind_rejected_at_boundary(self):
|
||||
from fastvideo.api.compat import _validate_continuation_state
|
||||
with pytest.raises(ValueError, match="non-empty string"):
|
||||
_validate_continuation_state(
|
||||
ContinuationState(kind="", payload={}))
|
||||
|
||||
def test_output_return_state_flag(self):
|
||||
request = GenerationRequest(
|
||||
prompt="x",
|
||||
output=OutputConfig(return_state=True),
|
||||
)
|
||||
# The typed public surface exposes the flag directly.
|
||||
assert request.output.return_state is True
|
||||
@@ -0,0 +1,298 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""gpu_pool-style flat-kwarg integration tests.
|
||||
|
||||
Mirrors the ``load_kwargs`` dict that the FastVideo-internal
|
||||
``ui/ltx2-streaming/server/gpu_pool.py`` passes to
|
||||
``VideoGenerator.from_pretrained(**load_kwargs)`` and asserts that the
|
||||
public typed ``GeneratorConfig`` surface (introduced across PRs 0-6)
|
||||
can represent it end-to-end, with no fields silently falling through
|
||||
to ``pipeline.experimental``.
|
||||
|
||||
This is the parity guard PR 7.6 depends on: the public gpu_pool
|
||||
upstream must be able to construct a typed ``GeneratorConfig`` without
|
||||
knowing any legacy LTX-2 kwarg name, and downstream Dynamo
|
||||
(``FastVideoArgGroup``) must be able to do the same.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.compat import (
|
||||
generator_config_to_fastvideo_args,
|
||||
legacy_from_pretrained_to_config,
|
||||
)
|
||||
|
||||
|
||||
# Mirrors FastVideo-internal/ui/ltx2-streaming/server/gpu_pool.py
|
||||
# :lines 233-260 (load_kwargs constructed for VideoGenerator.from_pretrained).
|
||||
#
|
||||
# One item from gpu_pool.py's load_kwargs is deliberately excluded:
|
||||
# - ``pipeline_config=<PipelineConfig instance>`` — an opaque Python
|
||||
# object; internal mutates it in place (``dit_config.quant_config =
|
||||
# FP4Config()``). The typed path for quantization is tracked in
|
||||
# "Known Technical Debt" in PR plan.md; ``pipeline_config`` as an
|
||||
# instance legitimately belongs in ``pipeline.experimental``.
|
||||
#
|
||||
# ``enable_torch_compile_text_encoder`` IS included below: its typed
|
||||
# home is ``CompileConfig.text_encoder_enabled`` (added post-review).
|
||||
# The legacy ``FastVideoArgs`` path does not yet consume it; the
|
||||
# realtime runtime (PR 7.6) reads it off the kwargs dict before
|
||||
# FastVideoArgs filtering.
|
||||
GPU_POOL_LOAD_KWARGS = {
|
||||
"config_model_path": "/models/ltx2-distilled/config",
|
||||
"num_gpus": 1,
|
||||
"dit_layerwise_offload": False,
|
||||
"use_fsdp_inference": False,
|
||||
"dit_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": True,
|
||||
"ltx2_vae_tiling": False,
|
||||
"ltx2_refine_enabled": True,
|
||||
"ltx2_refine_upsampler_path": "/models/ltx2-distilled/spatial_upsampler",
|
||||
"ltx2_refine_lora_path": "",
|
||||
"ltx2_refine_num_inference_steps": 2,
|
||||
"ltx2_refine_guidance_scale": 1.0,
|
||||
"ltx2_refine_add_noise": True,
|
||||
"enable_torch_compile": True,
|
||||
"enable_torch_compile_text_encoder": True,
|
||||
"torch_compile_kwargs": {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "max-autotune-no-cudagraphs",
|
||||
"dynamic": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class TestGpuPoolForwardTranslation:
|
||||
"""gpu_pool flat kwargs -> typed GeneratorConfig."""
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def config(self):
|
||||
return legacy_from_pretrained_to_config(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
GPU_POOL_LOAD_KWARGS,
|
||||
)
|
||||
|
||||
def test_model_path_set(self, config) -> None:
|
||||
assert config.model_path == "FastVideo/LTX2-Distilled-Diffusers"
|
||||
|
||||
def test_engine_basics(self, config) -> None:
|
||||
assert config.engine.num_gpus == 1
|
||||
assert config.engine.use_fsdp_inference is False
|
||||
|
||||
def test_offload_config(self, config) -> None:
|
||||
assert config.engine.offload.dit is False
|
||||
assert config.engine.offload.dit_layerwise is False
|
||||
assert config.engine.offload.vae is False
|
||||
assert config.engine.offload.text_encoder is False
|
||||
assert config.engine.offload.pin_cpu_memory is True
|
||||
|
||||
def test_compile_config_typed_fields_extracted(self, config) -> None:
|
||||
compile_config = config.engine.compile
|
||||
assert compile_config.enabled is True
|
||||
assert compile_config.text_encoder_enabled is True
|
||||
assert compile_config.backend == "inductor"
|
||||
assert compile_config.fullgraph is True
|
||||
assert compile_config.mode == "max-autotune-no-cudagraphs"
|
||||
assert compile_config.dynamic is False
|
||||
assert compile_config.extras == {}
|
||||
|
||||
def test_vae_tiling_routed_to_pipeline(self, config) -> None:
|
||||
assert config.pipeline.vae_tiling is False
|
||||
|
||||
def test_config_model_path_routed_to_components(self, config) -> None:
|
||||
assert config.pipeline.components.config_root == "/models/ltx2-distilled/config"
|
||||
|
||||
def test_refine_upsampler_routed_to_components(self, config) -> None:
|
||||
assert config.pipeline.components.upsampler_weights == (
|
||||
"/models/ltx2-distilled/spatial_upsampler")
|
||||
|
||||
def test_empty_refine_lora_becomes_none(self, config) -> None:
|
||||
# gpu_pool passes "" to keep refine LoRA disabled; typed schema
|
||||
# treats that as "no LoRA" rather than an empty-string path.
|
||||
assert config.pipeline.components.lora_path is None
|
||||
|
||||
def test_refine_preset_overrides(self, config) -> None:
|
||||
refine = config.pipeline.preset_overrides.get("refine", {})
|
||||
assert refine == {
|
||||
"enabled": True,
|
||||
"num_inference_steps": 2,
|
||||
"guidance_scale": 1.0,
|
||||
"add_noise": True,
|
||||
}
|
||||
|
||||
def test_no_experimental_leakage(self, config) -> None:
|
||||
"""Every gpu_pool kwarg should have a typed home — nothing should
|
||||
silently fall through to ``pipeline.experimental``."""
|
||||
assert config.pipeline.experimental == {}
|
||||
|
||||
|
||||
class TestGpuPoolReverseTranslation:
|
||||
"""typed GeneratorConfig -> FastVideoArgs kwargs reproduces the
|
||||
original gpu_pool flat-kwarg shape.
|
||||
|
||||
This is what lets PR 7.6 wire the public ``gpu_pool`` through
|
||||
``generator_config_to_fastvideo_args`` without the runtime noticing.
|
||||
"""
|
||||
|
||||
@pytest.fixture
|
||||
def args_kwargs(self, monkeypatch):
|
||||
from fastvideo import fastvideo_args as fva
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _capture(**kw):
|
||||
captured.update(kw)
|
||||
return _Captured(**kw)
|
||||
|
||||
class _Captured:
|
||||
|
||||
def __init__(self, **kw):
|
||||
self.kwargs = kw
|
||||
|
||||
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
|
||||
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
GPU_POOL_LOAD_KWARGS,
|
||||
)
|
||||
generator_config_to_fastvideo_args(config)
|
||||
return captured
|
||||
|
||||
def test_ltx2_refine_flags_reemitted(self, args_kwargs) -> None:
|
||||
assert args_kwargs["ltx2_refine_enabled"] is True
|
||||
assert args_kwargs["ltx2_refine_add_noise"] is True
|
||||
assert args_kwargs["ltx2_refine_num_inference_steps"] == 2
|
||||
assert args_kwargs["ltx2_refine_guidance_scale"] == 1.0
|
||||
|
||||
def test_refine_upsampler_path_reemitted(self, args_kwargs) -> None:
|
||||
assert args_kwargs["ltx2_refine_upsampler_path"] == (
|
||||
"/models/ltx2-distilled/spatial_upsampler")
|
||||
|
||||
def test_config_model_path_reemitted(self, args_kwargs) -> None:
|
||||
assert args_kwargs["config_model_path"] == "/models/ltx2-distilled/config"
|
||||
|
||||
def test_torch_compile_kwargs_reassembled(self, args_kwargs) -> None:
|
||||
assert args_kwargs["torch_compile_kwargs"] == {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
"mode": "max-autotune-no-cudagraphs",
|
||||
"dynamic": False,
|
||||
}
|
||||
|
||||
def test_vae_tiling_reemitted_with_legacy_name(self, args_kwargs) -> None:
|
||||
assert args_kwargs["ltx2_vae_tiling"] is False
|
||||
|
||||
def test_text_encoder_compile_reemitted(self, args_kwargs) -> None:
|
||||
# Present in the captured kwargs dict even though
|
||||
# ``FastVideoArgs.from_kwargs`` will filter it out — realtime
|
||||
# runtime upstream (PR 7.6) reads it off this dict.
|
||||
assert args_kwargs["enable_torch_compile_text_encoder"] is True
|
||||
|
||||
def test_no_stray_refine_dict(self, args_kwargs) -> None:
|
||||
"""preset_overrides.refine must flatten to ltx2_refine_* kwargs
|
||||
rather than landing as a nested ``refine`` kwarg that
|
||||
FastVideoArgs doesn't understand."""
|
||||
assert "refine" not in args_kwargs
|
||||
|
||||
|
||||
class TestRefineFlattenCoversAllTypedFields:
|
||||
"""Every field on LTX2Refine{Preset,Stage}Override must survive the
|
||||
round-trip through preset_overrides.refine back to ltx2_refine_*
|
||||
kwargs. Guards against the hardcoded-key-tuple regression where
|
||||
image_crf / video_position_offset_sec silently dropped."""
|
||||
|
||||
def test_all_fields_reemitted(self, monkeypatch) -> None:
|
||||
from fastvideo import fastvideo_args as fva
|
||||
from fastvideo.api.compat import (
|
||||
generator_config_to_fastvideo_args,
|
||||
)
|
||||
from fastvideo.api.schema import GeneratorConfig, PipelineSelection
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
class _Captured:
|
||||
|
||||
def __init__(self, **kw):
|
||||
self.kwargs = kw
|
||||
|
||||
def _capture(**kw):
|
||||
captured.update(kw)
|
||||
return _Captured(**kw)
|
||||
|
||||
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
|
||||
|
||||
refine_payload = {
|
||||
# Preset-override fields.
|
||||
"enabled": True,
|
||||
"add_noise": False,
|
||||
# Stage-override fields.
|
||||
"num_inference_steps": 3,
|
||||
"guidance_scale": 1.5,
|
||||
"image_crf": 18,
|
||||
"video_position_offset_sec": 2.5,
|
||||
}
|
||||
all_fields = (refine_preset_override_fields()
|
||||
| refine_stage_override_fields())
|
||||
assert set(refine_payload) == all_fields, (
|
||||
"payload must cover every typed field to exercise the flatten loop")
|
||||
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
pipeline=PipelineSelection(preset_overrides={"refine": refine_payload}),
|
||||
)
|
||||
generator_config_to_fastvideo_args(config)
|
||||
|
||||
for key, value in refine_payload.items():
|
||||
assert captured[f"ltx2_refine_{key}"] == value
|
||||
|
||||
|
||||
class TestCompileExtrasPreserved:
|
||||
"""Additional torch.compile kwargs beyond the four typed fields
|
||||
round-trip through ``CompileConfig.extras``."""
|
||||
|
||||
def test_extras_preserved(self, monkeypatch) -> None:
|
||||
from fastvideo import fastvideo_args as fva
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _capture(**kw):
|
||||
captured.update(kw)
|
||||
|
||||
class _Captured:
|
||||
|
||||
def __init__(self, **kw):
|
||||
self.kwargs = kw
|
||||
|
||||
return _Captured(**kw)
|
||||
|
||||
monkeypatch.setattr(fva.FastVideoArgs, "from_kwargs", _capture)
|
||||
|
||||
kwargs = deepcopy(GPU_POOL_LOAD_KWARGS)
|
||||
kwargs["torch_compile_kwargs"] = {
|
||||
"backend": "inductor",
|
||||
"options": {"triton.cudagraphs": False},
|
||||
"disable": False,
|
||||
}
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"FastVideo/LTX2-Distilled-Diffusers", kwargs)
|
||||
assert config.engine.compile.backend == "inductor"
|
||||
assert config.engine.compile.extras == {
|
||||
"options": {"triton.cudagraphs": False},
|
||||
"disable": False,
|
||||
}
|
||||
|
||||
generator_config_to_fastvideo_args(config)
|
||||
assert captured["torch_compile_kwargs"] == {
|
||||
"backend": "inductor",
|
||||
"options": {"triton.cudagraphs": False},
|
||||
"disable": False,
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for typed LTX-2 stage override dataclasses."""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.presets import get_preset, validate_stage_overrides
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
LTX2RefinePresetOverride,
|
||||
LTX2RefineStageOverride,
|
||||
refine_override_to_dict,
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
|
||||
|
||||
class TestRefineStageOverrideDataclass:
|
||||
|
||||
def test_all_fields_default_to_none(self) -> None:
|
||||
override = LTX2RefineStageOverride()
|
||||
assert override.num_inference_steps is None
|
||||
assert override.guidance_scale is None
|
||||
assert override.image_crf is None
|
||||
assert override.video_position_offset_sec is None
|
||||
|
||||
def test_explicit_construction(self) -> None:
|
||||
override = LTX2RefineStageOverride(
|
||||
num_inference_steps=2,
|
||||
guidance_scale=1.0,
|
||||
image_crf=18,
|
||||
video_position_offset_sec=2.5,
|
||||
)
|
||||
assert override.num_inference_steps == 2
|
||||
assert override.guidance_scale == 1.0
|
||||
assert override.image_crf == 18
|
||||
assert override.video_position_offset_sec == 2.5
|
||||
|
||||
def test_to_dict_drops_none(self) -> None:
|
||||
override = LTX2RefineStageOverride(num_inference_steps=3)
|
||||
assert refine_override_to_dict(override) == {
|
||||
"num_inference_steps": 3,
|
||||
}
|
||||
|
||||
def test_to_dict_with_all_fields(self) -> None:
|
||||
override = LTX2RefineStageOverride(
|
||||
num_inference_steps=2,
|
||||
guidance_scale=1.0,
|
||||
image_crf=18,
|
||||
video_position_offset_sec=0.0,
|
||||
)
|
||||
assert refine_override_to_dict(override) == {
|
||||
"num_inference_steps": 2,
|
||||
"guidance_scale": 1.0,
|
||||
"image_crf": 18,
|
||||
"video_position_offset_sec": 0.0,
|
||||
}
|
||||
|
||||
def test_fields_accessor_matches_dataclass(self) -> None:
|
||||
assert refine_stage_override_fields() == frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
"image_crf",
|
||||
"video_position_offset_sec",
|
||||
})
|
||||
|
||||
|
||||
class TestRefinePresetOverrideDataclass:
|
||||
|
||||
def test_all_fields_default_to_none(self) -> None:
|
||||
override = LTX2RefinePresetOverride()
|
||||
assert override.enabled is None
|
||||
assert override.add_noise is None
|
||||
|
||||
def test_to_dict_drops_none(self) -> None:
|
||||
override = LTX2RefinePresetOverride(enabled=True)
|
||||
assert refine_override_to_dict(override) == {
|
||||
"enabled": True,
|
||||
}
|
||||
|
||||
def test_to_dict_with_all_fields(self) -> None:
|
||||
override = LTX2RefinePresetOverride(enabled=True, add_noise=False)
|
||||
assert refine_override_to_dict(override) == {
|
||||
"enabled": True,
|
||||
"add_noise": False,
|
||||
}
|
||||
|
||||
def test_fields_accessor_matches_dataclass(self) -> None:
|
||||
assert refine_preset_override_fields() == frozenset({
|
||||
"enabled",
|
||||
"add_noise",
|
||||
})
|
||||
|
||||
|
||||
class TestStageOverridesMirrorPresetSchema:
|
||||
"""The ltx2_two_stage preset's refine stage schema must list
|
||||
exactly the :class:`LTX2RefineStageOverride` field names."""
|
||||
|
||||
def test_allowed_overrides_mirror_dataclass(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
preset = get_preset("ltx2_two_stage", "ltx2")
|
||||
refine_schema = next(
|
||||
s for s in preset.stage_schemas if s.name == "refine")
|
||||
assert refine_schema.allowed_overrides == refine_stage_override_fields()
|
||||
|
||||
def test_roundtrip_through_validate_stage_overrides(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
preset = get_preset("ltx2_two_stage", "ltx2")
|
||||
override = LTX2RefineStageOverride(
|
||||
num_inference_steps=3,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
validate_stage_overrides(
|
||||
preset, {"refine": refine_override_to_dict(override)})
|
||||
|
||||
def test_unknown_field_rejected(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
preset = get_preset("ltx2_two_stage", "ltx2")
|
||||
with pytest.raises(ConfigValidationError):
|
||||
validate_stage_overrides(
|
||||
preset, {"refine": {"unknown_key": 1}})
|
||||
@@ -111,7 +111,15 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
|
||||
"vae": True,
|
||||
"pin_cpu_memory": True,
|
||||
},
|
||||
"compile": {"enabled": False, "kwargs": {}},
|
||||
"compile": {
|
||||
"enabled": False,
|
||||
"text_encoder_enabled": None,
|
||||
"backend": None,
|
||||
"fullgraph": None,
|
||||
"mode": None,
|
||||
"dynamic": None,
|
||||
"extras": {},
|
||||
},
|
||||
"enable_stage_verification": True,
|
||||
"use_fsdp_inference": False,
|
||||
"disable_autocast": False,
|
||||
@@ -133,6 +141,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
|
||||
"override_pipeline_cls_name": None,
|
||||
"override_transformer_cls_name": None,
|
||||
},
|
||||
"vae_tiling": None,
|
||||
"preset_overrides": {},
|
||||
"experimental": {},
|
||||
},
|
||||
|
||||
@@ -334,7 +334,7 @@ class TestLtx2Presets:
|
||||
import fastvideo.registry # noqa: F401
|
||||
presets = get_presets_for_family("ltx2")
|
||||
names = {p.name for p in presets}
|
||||
assert names == {"ltx2_base", "ltx2_distilled"}
|
||||
assert names == {"ltx2_base", "ltx2_distilled", "ltx2_two_stage"}
|
||||
|
||||
def test_ltx2_base_lookup(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
@@ -349,6 +349,43 @@ class TestLtx2Presets:
|
||||
assert p.defaults["num_inference_steps"] == 8
|
||||
assert p.defaults["guidance_scale"] == 1.0
|
||||
|
||||
def test_ltx2_two_stage_is_two_stage(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
p = get_preset("ltx2_two_stage", "ltx2")
|
||||
assert len(p.stage_schemas) == 2
|
||||
assert p.stage_schemas[0].name == "denoise"
|
||||
assert p.stage_schemas[1].name == "refine"
|
||||
assert p.stage_schemas[1].kind == "refinement"
|
||||
|
||||
def test_ltx2_two_stage_stage_defaults(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
p = get_preset("ltx2_two_stage", "ltx2")
|
||||
refine = p.stage_defaults["refine"]
|
||||
# stage-2 refine only supports 2 or 3 denoising steps; preset
|
||||
# defaults to 2 (matches gpu_pool.py load_kwargs).
|
||||
assert refine["num_inference_steps"] == 2
|
||||
assert refine["guidance_scale"] == 1.0
|
||||
|
||||
def test_ltx2_two_stage_refine_overrides_valid(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
p = get_preset("ltx2_two_stage", "ltx2")
|
||||
validate_stage_overrides(
|
||||
p, {"refine": {"num_inference_steps": 3}})
|
||||
validate_stage_overrides(
|
||||
p, {"refine": {"guidance_scale": 1.0}})
|
||||
validate_stage_overrides(
|
||||
p, {"refine": {"image_crf": 18}})
|
||||
validate_stage_overrides(
|
||||
p, {"refine": {"video_position_offset_sec": 2.5}})
|
||||
|
||||
def test_ltx2_two_stage_rejects_unknown_refine_override(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
p = get_preset("ltx2_two_stage", "ltx2")
|
||||
with pytest.raises(ConfigValidationError):
|
||||
validate_stage_overrides(
|
||||
p, {"refine": {"bogus_field": 1}})
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Hunyuan preset integration
|
||||
@@ -548,3 +585,40 @@ class TestPresetCountIntegrity:
|
||||
import fastvideo.registry # noqa: F401
|
||||
names = get_all_preset_names()
|
||||
assert len(names) >= 37
|
||||
|
||||
|
||||
class TestPresetDefaultTypes:
|
||||
"""Preset ``defaults`` values must match the types on
|
||||
:class:`SamplingParam`. Assigning ``None`` to a typed-``str`` field
|
||||
(e.g. ``negative_prompt``) breaks downstream stages that assert the
|
||||
runtime type — see the CFG branch in
|
||||
``pipelines/stages/text_encoding.py:81``."""
|
||||
|
||||
def test_ltx2_cfg_defaults_are_off(self) -> None:
|
||||
"""SamplingParam's LTX-2 CFG class defaults must be 1.0 (CFG
|
||||
off). ``ForwardBatch.__post_init__`` force-enables
|
||||
``do_classifier_free_guidance`` when either
|
||||
``ltx2_cfg_scale_video`` or ``ltx2_cfg_scale_audio`` is != 1.0,
|
||||
so any non-1.0 default silently forces CFG on for every model
|
||||
family that doesn't explicitly override these fields. Guard
|
||||
against the regression that surfaced as the TurboDiffusion I2V
|
||||
SSIM crash (``text_encoding.py:81`` assertion on
|
||||
``negative_prompt``)."""
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
sp = SamplingParam()
|
||||
assert sp.ltx2_cfg_scale_video == 1.0
|
||||
assert sp.ltx2_cfg_scale_audio == 1.0
|
||||
|
||||
def test_no_preset_sets_negative_prompt_to_none(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
from fastvideo.api.presets import _PRESET_REGISTRY
|
||||
offenders = [
|
||||
f"{preset.model_family}/{preset.name}"
|
||||
for preset in _PRESET_REGISTRY.values()
|
||||
if preset.defaults.get("negative_prompt", "") is None
|
||||
]
|
||||
assert not offenders, (
|
||||
"These presets set negative_prompt=None, which violates "
|
||||
"SamplingParam.negative_prompt's typed str contract and "
|
||||
"crashes the CFG path in text_encoding. Use \"\" instead:\n"
|
||||
+ "\n".join(f" - {p}" for p in offenders))
|
||||
|
||||
@@ -46,24 +46,42 @@ def _flatten_status_section(section: dict, valid_statuses: set[str]) -> set[str]
|
||||
return names
|
||||
|
||||
|
||||
def _get_extra_dataclass_fields(package_name: str, base_cls: type) -> set[str]:
|
||||
package = importlib.import_module(package_name)
|
||||
def _get_extra_dataclass_fields(
|
||||
package_names: str | tuple[str, ...],
|
||||
base_cls: type,
|
||||
) -> set[str]:
|
||||
"""Collect dataclass fields declared on ``base_cls`` subclasses found
|
||||
under any of the given package roots.
|
||||
|
||||
Accepts either a single package name (string) or a tuple of package
|
||||
roots — the latter supports the PR 6 colocation where each model
|
||||
family's ``PipelineConfig`` subclass moves from
|
||||
``fastvideo.configs.pipelines.<family>`` to
|
||||
``fastvideo.pipelines.basic.<family>.pipeline_configs``.
|
||||
"""
|
||||
if isinstance(package_names, str):
|
||||
package_names = (package_names, )
|
||||
|
||||
base_fields = {f.name for f in dataclasses.fields(base_cls)}
|
||||
extras: set[str] = set()
|
||||
if not hasattr(package, "__path__"):
|
||||
return extras
|
||||
for _, modname, _ in pkgutil.iter_modules(package.__path__):
|
||||
if modname == "__pycache__":
|
||||
|
||||
for package_name in package_names:
|
||||
package = importlib.import_module(package_name)
|
||||
if not hasattr(package, "__path__"):
|
||||
continue
|
||||
module = importlib.import_module(f"{package_name}.{modname}")
|
||||
for obj in vars(module).values():
|
||||
if (
|
||||
isinstance(obj, type)
|
||||
and dataclasses.is_dataclass(obj)
|
||||
and issubclass(obj, base_cls)
|
||||
and obj is not base_cls
|
||||
):
|
||||
extras.update(f.name for f in dataclasses.fields(obj) if f.name not in base_fields)
|
||||
for _, modname, is_pkg in pkgutil.walk_packages(
|
||||
package.__path__, prefix=f"{package_name}."):
|
||||
if modname.endswith(".__pycache__"):
|
||||
continue
|
||||
module = importlib.import_module(modname)
|
||||
for obj in vars(module).values():
|
||||
if (isinstance(obj, type)
|
||||
and dataclasses.is_dataclass(obj)
|
||||
and issubclass(obj, base_cls)
|
||||
and obj is not base_cls):
|
||||
extras.update(
|
||||
f.name for f in dataclasses.fields(obj)
|
||||
if f.name not in base_fields)
|
||||
return extras
|
||||
|
||||
|
||||
@@ -177,7 +195,10 @@ def test_pipeline_config_base_fields_are_classified() -> None:
|
||||
|
||||
def test_pipeline_config_extension_fields_are_classified() -> None:
|
||||
inventory = _load_inventory()
|
||||
expected = _get_extra_dataclass_fields("fastvideo.configs.pipelines", PipelineConfig)
|
||||
expected = _get_extra_dataclass_fields(
|
||||
("fastvideo.configs.pipelines", "fastvideo.pipelines.basic"),
|
||||
PipelineConfig,
|
||||
)
|
||||
actual = _flatten_status_section(
|
||||
inventory["surfaces"]["pipeline_config_extensions"],
|
||||
set(inventory["status_definitions"]),
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,154 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Protocol schema tests for the streaming server.
|
||||
|
||||
Covers:
|
||||
|
||||
* accepted client messages parse into the correct discriminated model
|
||||
* unknown ``type`` values raise validation errors
|
||||
* server-side messages serialize to the expected wire shape
|
||||
* continuation_state field on session_init_v2 carries through
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
ContinuationStateSnapshot,
|
||||
ErrorMessage,
|
||||
GpuAssigned,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
|
||||
|
||||
class TestClientMessageParsing:
|
||||
|
||||
def test_session_init_v2_minimal(self):
|
||||
parsed = parse_client_message({"type": "session_init_v2"})
|
||||
assert isinstance(parsed, SessionInitV2)
|
||||
assert parsed.curated_prompts == []
|
||||
assert parsed.stream_mode == "av_fmp4"
|
||||
|
||||
def test_session_init_v2_full(self):
|
||||
raw = {
|
||||
"type": "session_init_v2",
|
||||
"client_id": "client-1",
|
||||
"preset": "ltx2_two_stage",
|
||||
"preset_label": "2x refine",
|
||||
"curated_prompts": ["a fox", "a deer"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
"single_clip_mode": True,
|
||||
"stream_mode": "av_fmp4",
|
||||
"continuation_state": {
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, "segment_index": 2},
|
||||
},
|
||||
}
|
||||
parsed = parse_client_message(raw)
|
||||
assert isinstance(parsed, SessionInitV2)
|
||||
assert parsed.preset == "ltx2_two_stage"
|
||||
assert parsed.curated_prompts == ["a fox", "a deer"]
|
||||
assert parsed.continuation_state["kind"] == "ltx2.v1"
|
||||
|
||||
def test_segment_prompt_source(self):
|
||||
parsed = parse_client_message({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "hello world",
|
||||
"source": "curated",
|
||||
"seed": 7,
|
||||
})
|
||||
assert isinstance(parsed, SegmentPromptSource)
|
||||
assert parsed.source == "curated"
|
||||
assert parsed.seed == 7
|
||||
|
||||
def test_snapshot_state(self):
|
||||
parsed = parse_client_message({"type": "snapshot_state"})
|
||||
assert isinstance(parsed, SnapshotState)
|
||||
|
||||
def test_unknown_type_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"type": "not_a_real_message"})
|
||||
|
||||
def test_missing_type_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"prompt": "x"})
|
||||
|
||||
def test_segment_prompt_source_requires_prompt(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"type": "segment_prompt_source"})
|
||||
|
||||
|
||||
class TestServerMessageSerialization:
|
||||
|
||||
def test_queue_status(self):
|
||||
msg = QueueStatus(position=3, queue_depth=5)
|
||||
assert msg.model_dump() == {
|
||||
"type": "queue_status",
|
||||
"position": 3,
|
||||
"queue_depth": 5,
|
||||
}
|
||||
|
||||
def test_gpu_assigned(self):
|
||||
msg = GpuAssigned(gpu_id=1, session_timeout=300)
|
||||
assert msg.model_dump()["type"] == "gpu_assigned"
|
||||
|
||||
def test_ltx2_stream_start(self):
|
||||
msg = Ltx2StreamStart(
|
||||
preset="ltx2_two_stage",
|
||||
width=1024, height=1536, fps=24, num_frames=121,
|
||||
)
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["type"] == "ltx2_stream_start"
|
||||
assert dumped["width"] == 1024
|
||||
|
||||
def test_ltx2_segment_start(self):
|
||||
msg = Ltx2SegmentStart(
|
||||
segment_idx=0,
|
||||
prompt="a fox",
|
||||
total_steps=8,
|
||||
)
|
||||
assert msg.model_dump()["segment_idx"] == 0
|
||||
|
||||
def test_step_complete(self):
|
||||
msg = StepComplete(segment_idx=0, step=1, total_steps=8)
|
||||
assert msg.model_dump()["stage"] == "denoise"
|
||||
|
||||
def test_media_init_has_mode(self):
|
||||
msg = MediaInit(segment_idx=0, stream_id="abc")
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["mode"] == "av_fmp4"
|
||||
assert "avc1" in dumped["mime"]
|
||||
|
||||
def test_media_segment_complete(self):
|
||||
msg = MediaSegmentComplete(
|
||||
segment_idx=0, stream_id="abc", chunks=4,
|
||||
)
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["chunks"] == 4
|
||||
|
||||
def test_ltx2_segment_complete(self):
|
||||
msg = Ltx2SegmentComplete(segment_idx=0, generation_time_ms=1234.5)
|
||||
assert msg.model_dump()["generation_time_ms"] == 1234.5
|
||||
|
||||
def test_error_message_code_restricted(self):
|
||||
with pytest.raises(ValidationError):
|
||||
ErrorMessage(code="not_a_code", message="x")
|
||||
|
||||
def test_continuation_state_snapshot(self):
|
||||
msg = ContinuationStateSnapshot(state={
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1},
|
||||
})
|
||||
assert msg.model_dump()["state"]["kind"] == "ltx2.v1"
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""End-to-end WebSocket smoke for the streaming server skeleton.
|
||||
|
||||
Uses a mock generator so these tests run CPU-only (no GPU, no model
|
||||
weights). Skips the fMP4 assertions when ``ffmpeg`` is missing.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("starlette")
|
||||
from starlette.testclient import TestClient # noqa: E402
|
||||
|
||||
from fastvideo.api.schema import ( # noqa: E402
|
||||
ContinuationState,
|
||||
GeneratorConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
StreamingConfig,
|
||||
GenerationRequest,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.server import build_app # noqa: E402
|
||||
|
||||
|
||||
_FFMPEG_AVAILABLE = shutil.which("ffmpeg") is not None
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MockGenerator:
|
||||
width: int = 64
|
||||
height: int = 64
|
||||
fps: int = 12
|
||||
num_frames: int = 12
|
||||
return_state: bool = True
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
frames = [
|
||||
np.full((self.height, self.width, 3), i * 5, dtype=np.uint8)
|
||||
for i in range(self.num_frames)
|
||||
]
|
||||
state = (ContinuationState(
|
||||
kind="ltx2.v1",
|
||||
payload={
|
||||
"schema_version": 1,
|
||||
"segment_index": 0,
|
||||
"source_prompt": request.prompt,
|
||||
},
|
||||
) if self.return_state else None)
|
||||
return {
|
||||
"frames": frames,
|
||||
"audio_sample_rate": 24000,
|
||||
"state": state,
|
||||
}
|
||||
|
||||
|
||||
def _build_serve_config() -> ServeConfig:
|
||||
return ServeConfig(
|
||||
generator=GeneratorConfig(model_path="/models/fake"),
|
||||
default_request=GenerationRequest(
|
||||
sampling=SamplingConfig(
|
||||
num_frames=12,
|
||||
height=64,
|
||||
width=64,
|
||||
fps=12,
|
||||
num_inference_steps=1,
|
||||
),
|
||||
),
|
||||
streaming=StreamingConfig(
|
||||
session_timeout_seconds=60,
|
||||
generation_segment_cap=2,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _build_client() -> tuple[TestClient, _MockGenerator]:
|
||||
generator = _MockGenerator()
|
||||
app = build_app(_build_serve_config(), generator)
|
||||
return TestClient(app), generator
|
||||
|
||||
|
||||
class TestHealth:
|
||||
|
||||
def test_health_endpoint_reports_stream_mode(self):
|
||||
client, _ = _build_client()
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["stream_mode"] == "av_fmp4"
|
||||
assert body["sessions"] == 0
|
||||
|
||||
|
||||
class TestSessionHandshake:
|
||||
|
||||
def test_rejects_non_session_init_opening_frame(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "segment_prompt_source", "prompt": "x"})
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
assert err["code"] == "invalid_message"
|
||||
|
||||
def test_rejects_unknown_message_on_init(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "not_a_message"})
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
|
||||
def test_emits_queue_and_gpu_assigned_on_valid_init(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({
|
||||
"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox"],
|
||||
})
|
||||
assert ws.receive_json()["type"] == "queue_status"
|
||||
assert ws.receive_json()["type"] == "gpu_assigned"
|
||||
assert ws.receive_json()["type"] == "ltx2_stream_start"
|
||||
|
||||
def test_init_hydrates_continuation_state(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({
|
||||
"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"continuation_state": {
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, "segment_index": 3},
|
||||
},
|
||||
})
|
||||
# Drain handshake frames
|
||||
ws.receive_json() # queue_status
|
||||
ws.receive_json() # gpu_assigned
|
||||
ws.receive_json() # ltx2_stream_start
|
||||
# Ask the server for the state back; it should echo what we sent.
|
||||
ws.send_json({"type": "snapshot_state"})
|
||||
snap = ws.receive_json()
|
||||
assert snap["type"] == "continuation_state_snapshot"
|
||||
assert snap["state"]["kind"] == "ltx2.v1"
|
||||
assert snap["state"]["payload"]["segment_index"] == 3
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _FFMPEG_AVAILABLE, reason="ffmpeg not installed")
|
||||
class TestSegmentFlow:
|
||||
|
||||
def test_segment_generates_media_init_plus_complete(self):
|
||||
client, generator = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage"})
|
||||
for _ in range(3):
|
||||
ws.receive_json() # queue_status + gpu_assigned + stream_start
|
||||
|
||||
ws.send_json({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "a test segment",
|
||||
"num_inference_steps": 1,
|
||||
})
|
||||
start = ws.receive_json()
|
||||
assert start["type"] == "ltx2_segment_start"
|
||||
assert start["segment_idx"] == 0
|
||||
step = ws.receive_json()
|
||||
assert step["type"] == "step_complete"
|
||||
media_init = ws.receive_json()
|
||||
assert media_init["type"] == "media_init"
|
||||
# Then one or more binary frames until media_segment_complete.
|
||||
saw_binary = False
|
||||
while True:
|
||||
msg = ws.receive()
|
||||
if "bytes" in msg and msg["bytes"]:
|
||||
saw_binary = True
|
||||
continue
|
||||
parsed = _as_json(msg)
|
||||
if parsed is None:
|
||||
continue
|
||||
if parsed["type"] == "media_segment_complete":
|
||||
break
|
||||
assert saw_binary
|
||||
final = ws.receive_json()
|
||||
assert final["type"] == "ltx2_segment_complete"
|
||||
assert final["segment_idx"] == 0
|
||||
|
||||
|
||||
class TestContinuationStatePersistence:
|
||||
|
||||
def test_snapshot_after_segment_carries_generator_state(self):
|
||||
if not _FFMPEG_AVAILABLE:
|
||||
pytest.skip("ffmpeg not installed")
|
||||
client, generator = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage"})
|
||||
for _ in range(3):
|
||||
ws.receive_json()
|
||||
ws.send_json({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "a cat",
|
||||
"num_inference_steps": 1,
|
||||
})
|
||||
_drain_until(ws, "ltx2_segment_complete")
|
||||
ws.send_json({"type": "snapshot_state"})
|
||||
snap = ws.receive_json()
|
||||
assert snap["type"] == "continuation_state_snapshot"
|
||||
assert snap["state"]["kind"] == "ltx2.v1"
|
||||
assert snap["state"]["payload"]["source_prompt"] == "a cat"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
def _drain_until(ws, target_type: str) -> dict[str, Any]:
|
||||
while True:
|
||||
msg = ws.receive()
|
||||
if "text" in msg and msg["text"]:
|
||||
import json
|
||||
|
||||
parsed = json.loads(msg["text"])
|
||||
if parsed.get("type") == target_type:
|
||||
return parsed
|
||||
# skip binary / other
|
||||
|
||||
|
||||
def _as_json(msg: dict[str, Any]) -> dict[str, Any] | None:
|
||||
if "text" not in msg or not msg["text"]:
|
||||
return None
|
||||
import json
|
||||
|
||||
return json.loads(msg["text"])
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Session lifecycle tests."""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
|
||||
|
||||
class TestSessionStateMachine:
|
||||
|
||||
def test_starts_initializing(self):
|
||||
s = Session()
|
||||
assert s.state is SessionState.INITIALIZING
|
||||
|
||||
def test_legal_sequence(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.COMPLETE)
|
||||
assert s.state is SessionState.COMPLETE
|
||||
|
||||
def test_active_self_loop_allowed(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.ACTIVE) # re-asserting is fine
|
||||
assert s.state is SessionState.ACTIVE
|
||||
|
||||
def test_illegal_backwards_transition(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.INITIALIZING)
|
||||
|
||||
def test_cannot_leave_terminal_state(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.COMPLETE)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.ACTIVE)
|
||||
|
||||
def test_error_terminal(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.ERROR)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.ACTIVE)
|
||||
|
||||
def test_transition_updates_activity(self):
|
||||
s = Session()
|
||||
prior = s.last_activity
|
||||
time.sleep(0.001)
|
||||
s.transition(SessionState.QUEUED)
|
||||
assert s.last_activity > prior
|
||||
|
||||
def test_segment_cap(self):
|
||||
s = Session()
|
||||
s.segment_idx = 5
|
||||
assert s.segment_cap_reached(5) is True
|
||||
assert s.segment_cap_reached(6) is False
|
||||
|
||||
|
||||
class TestSessionManager:
|
||||
|
||||
def test_create_assigns_unique_ids(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=2)
|
||||
a = mgr.create()
|
||||
b = mgr.create()
|
||||
assert a.id != b.id
|
||||
assert len(mgr) == 2
|
||||
|
||||
def test_max_sessions_enforced(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
|
||||
mgr.create()
|
||||
with pytest.raises(SessionRejected):
|
||||
mgr.create()
|
||||
|
||||
def test_close_releases_slot(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
|
||||
s = mgr.create()
|
||||
mgr.close(s.id)
|
||||
assert len(mgr) == 0
|
||||
# Now can create again.
|
||||
mgr.create()
|
||||
|
||||
def test_reap_timed_out_flags_idle_sessions(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
|
||||
s = mgr.create()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.last_activity = time.monotonic() - 10 # 10s ago, past the 1s budget
|
||||
dead = mgr.reap_timed_out()
|
||||
assert s.id in dead
|
||||
|
||||
def test_reap_skips_terminal_states(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
|
||||
s = mgr.create()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.ERROR)
|
||||
s.last_activity = time.monotonic() - 10
|
||||
assert s.id not in mgr.reap_timed_out()
|
||||
|
||||
def test_active_sessions_filter(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=4)
|
||||
a = mgr.create()
|
||||
a.transition(SessionState.QUEUED)
|
||||
a.transition(SessionState.GPU_BINDING)
|
||||
a.transition(SessionState.ACTIVE)
|
||||
b = mgr.create() # INITIALIZING
|
||||
assert mgr.active_sessions() == [a]
|
||||
assert b not in mgr.active_sessions()
|
||||
@@ -0,0 +1,90 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the session init-image persistence helper."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image,
|
||||
)
|
||||
|
||||
|
||||
def _png_bytes(size: tuple[int, int] = (64, 64)) -> bytes:
|
||||
buffer = io.BytesIO()
|
||||
Image.new("RGB", size, color=(10, 20, 30)).save(buffer, format="PNG")
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
class TestPersistSessionInitImage:
|
||||
|
||||
def test_none_payload_returns_none(self):
|
||||
assert persist_session_init_image(None) is None
|
||||
assert persist_session_init_image({}) is None
|
||||
|
||||
def test_non_object_payload_rejected(self):
|
||||
with pytest.raises(ValueError):
|
||||
persist_session_init_image("not-a-dict")
|
||||
|
||||
def test_png_payload_persists(self, tmp_path):
|
||||
data = _png_bytes()
|
||||
image = persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"name": "ref.png",
|
||||
"data": base64.b64encode(data).decode("ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
assert image is not None
|
||||
assert os.path.exists(image.path)
|
||||
assert image.mime == "image/png"
|
||||
assert image.path.endswith(".png")
|
||||
with open(image.path, "rb") as f:
|
||||
assert f.read() == data
|
||||
|
||||
def test_unknown_mime_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="mime"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/bmp",
|
||||
"data": "ignored",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_bad_base64_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="base64"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": "not!base64!",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_empty_data_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="empty"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": "",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_display_name_sanitized(self, tmp_path):
|
||||
image = persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"name": "../evil/../name.png",
|
||||
"data": base64.b64encode(_png_bytes()).decode("ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
assert image is not None
|
||||
assert image.display_name == "name.png"
|
||||
|
||||
def test_oversize_rejected(self, tmp_path):
|
||||
from fastvideo.entrypoints.streaming import session_init_image as mod
|
||||
|
||||
original = mod._MAX_IMAGE_BYTES
|
||||
mod._MAX_IMAGE_BYTES = 100
|
||||
try:
|
||||
with pytest.raises(ValueError, match="limit"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": base64.b64encode(_png_bytes((512, 512))).decode(
|
||||
"ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
finally:
|
||||
mod._MAX_IMAGE_BYTES = original
|
||||
@@ -0,0 +1,186 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the streaming SessionStore and BlobStore.
|
||||
|
||||
Covers:
|
||||
|
||||
* ``store`` / ``snapshot`` / ``drop`` lifecycle for the in-memory store
|
||||
* ``hydrate`` with and without an explicit session id
|
||||
* blob store insert / get / drop semantics
|
||||
* thread-safety under concurrent writes (smoke)
|
||||
* round-trip a LTX-2 continuation through snapshot + hydrate across a
|
||||
session boundary (the "export and resume" flow the PR plan calls out)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
BlobStore,
|
||||
InMemoryBlobStore,
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.pipelines.basic.ltx2.continuation import (
|
||||
LTX2_CONTINUATION_KIND,
|
||||
LTX2ContinuationState,
|
||||
)
|
||||
|
||||
|
||||
class TestInMemoryBlobStore:
|
||||
|
||||
def test_is_blob_store(self):
|
||||
assert isinstance(InMemoryBlobStore(), BlobStore)
|
||||
|
||||
def test_put_then_get_returns_same_bytes(self):
|
||||
store = InMemoryBlobStore()
|
||||
blob_id = store.put(b"hello")
|
||||
assert store.get(blob_id) == b"hello"
|
||||
|
||||
def test_put_returns_distinct_ids(self):
|
||||
store = InMemoryBlobStore()
|
||||
id_a = store.put(b"a")
|
||||
id_b = store.put(b"b")
|
||||
assert id_a != id_b
|
||||
|
||||
def test_get_missing_raises_keyerror(self):
|
||||
store = InMemoryBlobStore()
|
||||
with pytest.raises(KeyError):
|
||||
store.get("nonexistent")
|
||||
|
||||
def test_drop_removes_blob(self):
|
||||
store = InMemoryBlobStore()
|
||||
blob_id = store.put(b"payload")
|
||||
store.drop(blob_id)
|
||||
assert blob_id not in store
|
||||
with pytest.raises(KeyError):
|
||||
store.get(blob_id)
|
||||
|
||||
def test_drop_missing_is_noop(self):
|
||||
store = InMemoryBlobStore()
|
||||
store.drop("not-there") # no raise
|
||||
|
||||
def test_contains(self):
|
||||
store = InMemoryBlobStore()
|
||||
blob_id = store.put(b"x")
|
||||
assert blob_id in store
|
||||
assert "other" not in store
|
||||
|
||||
|
||||
class TestInMemorySessionStore:
|
||||
|
||||
def test_is_session_store(self):
|
||||
assert isinstance(InMemorySessionStore(), SessionStore)
|
||||
|
||||
def test_store_then_snapshot(self):
|
||||
store = InMemorySessionStore()
|
||||
state = ContinuationState(kind="ltx2.v1", payload={"x": 1})
|
||||
store.store("sess-1", state)
|
||||
assert store.snapshot("sess-1") is state
|
||||
|
||||
def test_snapshot_missing_returns_none(self):
|
||||
store = InMemorySessionStore()
|
||||
assert store.snapshot("missing") is None
|
||||
|
||||
def test_store_overwrites_prior_state(self):
|
||||
store = InMemorySessionStore()
|
||||
first = ContinuationState(kind="ltx2.v1", payload={"v": 1})
|
||||
second = ContinuationState(kind="ltx2.v1", payload={"v": 2})
|
||||
store.store("s", first)
|
||||
store.store("s", second)
|
||||
assert store.snapshot("s").payload["v"] == 2
|
||||
|
||||
def test_hydrate_assigns_new_session_id(self):
|
||||
store = InMemorySessionStore()
|
||||
state = ContinuationState(kind="ltx2.v1", payload={})
|
||||
sid = store.hydrate(state)
|
||||
assert sid
|
||||
assert store.snapshot(sid) is state
|
||||
|
||||
def test_hydrate_with_explicit_session_id(self):
|
||||
store = InMemorySessionStore()
|
||||
state = ContinuationState(kind="ltx2.v1", payload={})
|
||||
sid = store.hydrate(state, session_id="pinned-id")
|
||||
assert sid == "pinned-id"
|
||||
assert store.snapshot("pinned-id") is state
|
||||
|
||||
def test_drop_forgets_session(self):
|
||||
store = InMemorySessionStore()
|
||||
store.store("s", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
store.drop("s")
|
||||
assert store.snapshot("s") is None
|
||||
assert "s" not in store
|
||||
|
||||
def test_iter_yields_session_ids(self):
|
||||
store = InMemorySessionStore()
|
||||
store.store("a", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
store.store("b", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
assert sorted(store) == ["a", "b"]
|
||||
|
||||
def test_len(self):
|
||||
store = InMemorySessionStore()
|
||||
assert len(store) == 0
|
||||
store.store("x", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
assert len(store) == 1
|
||||
|
||||
def test_concurrent_store_is_safe(self):
|
||||
"""Smoke-check the lock: 200 parallel stores settle to 200 ids."""
|
||||
store = InMemorySessionStore()
|
||||
|
||||
def write(i: int) -> None:
|
||||
store.store(
|
||||
f"s-{i}",
|
||||
ContinuationState(kind="ltx2.v1", payload={"i": i}),
|
||||
)
|
||||
|
||||
threads = [threading.Thread(target=write, args=(i,)) for i in range(200)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
assert len(store) == 200
|
||||
|
||||
|
||||
class TestSnapshotHydrateRoundTrip:
|
||||
"""Session boundary: snapshot + hydrate preserves the full LTX-2 state."""
|
||||
|
||||
def test_end_to_end_ltx2_session_migration(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
sessions = InMemorySessionStore()
|
||||
|
||||
typed = LTX2ContinuationState(
|
||||
segment_index=4,
|
||||
video_frames=[
|
||||
np.full((32, 32, 3), i * 5, dtype=np.uint8) for i in range(3)
|
||||
],
|
||||
audio_latents=torch.randn(1, 4, 8, 32, dtype=torch.float32),
|
||||
audio_sample_rate=24000,
|
||||
audio_conditioning_num_frames=5,
|
||||
video_position_offset_sec=0.25,
|
||||
)
|
||||
envelope = typed.to_continuation_state(blob_store=blob_store)
|
||||
|
||||
sessions.store("session-a", envelope)
|
||||
snapshot = sessions.snapshot("session-a")
|
||||
assert snapshot is not None
|
||||
assert snapshot.kind == LTX2_CONTINUATION_KIND
|
||||
|
||||
# Simulate a migration: drop the first session, hydrate a new one
|
||||
# from the snapshot, and reconstruct the typed state.
|
||||
sessions.drop("session-a")
|
||||
new_sid = sessions.hydrate(snapshot)
|
||||
assert new_sid != "session-a"
|
||||
rebuilt = sessions.snapshot(new_sid)
|
||||
assert rebuilt is snapshot
|
||||
|
||||
restored = LTX2ContinuationState.from_continuation_state(
|
||||
rebuilt, blob_store=blob_store)
|
||||
assert restored.segment_index == typed.segment_index
|
||||
assert restored.audio_sample_rate == typed.audio_sample_rate
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, typed.audio_latents)
|
||||
assert len(restored.video_frames) == len(typed.video_frames)
|
||||
@@ -0,0 +1,99 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the fMP4 encoder.
|
||||
|
||||
These tests require ``ffmpeg`` on PATH. Skip when missing so the suite
|
||||
stays CPU/CI friendly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import shutil
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
shutil.which("ffmpeg") is None,
|
||||
reason="ffmpeg not installed",
|
||||
)
|
||||
|
||||
|
||||
def _frame(width: int, height: int, value: int = 128) -> np.ndarray:
|
||||
return np.full((height, width, 3), value, dtype=np.uint8)
|
||||
|
||||
|
||||
def test_encoder_emits_init_then_media_chunks():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
chunks: list[FragmentedMP4Chunk] = []
|
||||
async with enc:
|
||||
frames = [_frame(64, 64, v) for v in range(4, 28)]
|
||||
async for chunk in enc.encode(frames):
|
||||
chunks.append(chunk)
|
||||
assert len(chunks) > 0
|
||||
assert chunks[0].kind == "init"
|
||||
assert all(c.stream_id == enc.stream_id for c in chunks)
|
||||
assert all(c.segment_idx == 0 for c in chunks)
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_init_chunk_is_fmp4():
|
||||
"""The first chunk must contain the ``ftyp`` box (fMP4 init segment)."""
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
first_chunk = None
|
||||
async with enc:
|
||||
async for chunk in enc.encode([_frame(64, 64, 20)] * 24):
|
||||
first_chunk = chunk
|
||||
break
|
||||
assert first_chunk is not None
|
||||
assert first_chunk.kind == "init"
|
||||
# Box header: 4 bytes length, 4 bytes type. "ftyp" should appear
|
||||
# near the start of the init segment.
|
||||
assert b"ftyp" in first_chunk.data[:32]
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_rejects_non_ndarray_frames():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
async with enc:
|
||||
with pytest.raises(TypeError):
|
||||
async for _ in enc.encode(["not-a-frame"]):
|
||||
pass
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_rejects_wrong_shape():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
async with enc:
|
||||
with pytest.raises(ValueError):
|
||||
async for _ in enc.encode(
|
||||
[np.zeros((64, 64, 4), dtype=np.uint8)]):
|
||||
pass
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_close_is_idempotent():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
await enc.__aenter__()
|
||||
await enc.close()
|
||||
await enc.close() # no raise
|
||||
|
||||
asyncio.run(run())
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user