Compare commits
37
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 | ||
|
|
4ddcdf541f | ||
|
|
0e3529869c | ||
|
|
e1e0d91c00 | ||
|
|
145a3f166b | ||
|
|
88a5a933ab |
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
|
||||
|
||||
+1
-1
@@ -105,7 +105,7 @@ pull_request_rules:
|
||||
- files~=^fastvideo/pipelines/samplers/
|
||||
- files~=^fastvideo/entrypoints/
|
||||
- files~=^fastvideo/worker/
|
||||
- files~=^fastvideo/configs/sample/
|
||||
- files~=^fastvideo/api/sampling_param
|
||||
- files~=^fastvideo/configs/pipelines/
|
||||
- files~=^examples/inference/
|
||||
- -closed
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -62,9 +62,9 @@ This page contains the complete API reference for the FastVideo library.
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.configs.sample
|
||||
#### fastvideo.api.sampling_param
|
||||
|
||||
::: fastvideo.configs.sample
|
||||
::: fastvideo.api.sampling_param
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
|
||||
@@ -173,7 +173,7 @@ Applied by Mergify based on which paths you modified. Multiple scope labels can
|
||||
| Label | File paths that trigger it |
|
||||
|-------|---------------------------|
|
||||
| `scope: training` | `fastvideo/train/`, `fastvideo/training/`, `fastvideo/distillation/`, `examples/train/`, `examples/training/`, `examples/distill/` |
|
||||
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/configs/sample/`, `fastvideo/configs/pipelines/`, `examples/inference/` |
|
||||
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/api/sampling_param.py`, `fastvideo/configs/pipelines/`, `examples/inference/` |
|
||||
| `scope: attention` | `fastvideo/attention/` |
|
||||
| `scope: kernel` | `fastvideo-kernel/`, `csrc/` |
|
||||
| `scope: data` | `fastvideo/dataset/`, `fastvideo/pipelines/preprocess/`, `examples/preprocessing/` |
|
||||
|
||||
@@ -44,7 +44,7 @@ FastVideo maps a Diffusers-style repo into a pipeline like:
|
||||
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
|
||||
weight name translation.
|
||||
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
|
||||
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
|
||||
- `fastvideo/api/sampling_param.py`: runtime sampling parameters.
|
||||
- `fastvideo/pipelines/basic/*`: end-to-end pipeline logic built from stages.
|
||||
- `model_index.json`: the HF repo entrypoint that maps component names to
|
||||
classes and weight files.
|
||||
@@ -55,7 +55,7 @@ Minimal usage example (based on `examples/inference/basic/basic.py`):
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
|
||||
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
|
||||
@@ -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`).
|
||||
@@ -319,7 +321,8 @@ Purpose:
|
||||
|
||||
- `fastvideo/configs/pipelines/` describes pipeline wiring and model module
|
||||
names.
|
||||
- `fastvideo/configs/sample/` defines default runtime parameters.
|
||||
- `fastvideo/api/sampling_param.py` defines runtime sampling parameters.
|
||||
Defaults come from profiles in `fastvideo/pipelines/basic/<family>/profiles.py`.
|
||||
|
||||
Action:
|
||||
|
||||
@@ -347,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
|
||||
@@ -474,7 +478,7 @@ FastVideo integration.
|
||||
3. Pipeline wiring.
|
||||
- Pipeline: `fastvideo/pipelines/basic/wan/wan_pipeline.py`
|
||||
- Pipeline config: `fastvideo/configs/pipelines/wan.py`
|
||||
- Sampling defaults: `fastvideo/configs/sample/wan.py`
|
||||
- Sampling defaults: `fastvideo/pipelines/basic/wan/profiles.py`
|
||||
|
||||
4. Minimal example.
|
||||
- Script: `examples/inference/basic/basic.py`
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
status_definitions:
|
||||
kept: "Public field remains on a public adapter surface with the same meaning."
|
||||
moved: "Public field remains supported but normalizes into a different nested path."
|
||||
profile_owned: "Public field remains supported only through a model/profile-specific surface."
|
||||
preset_owned: "Public field remains supported only through a model/preset-specific surface."
|
||||
compatibility_only: "Legacy public field remains adapter-only during migration and is not part of the canonical typed schema."
|
||||
private_only: "Field should only be handled by private adapters and is not a public FastVideo compatibility promise."
|
||||
internal_only: "Field is runtime/config plumbing and should not be part of the new public typed inference API."
|
||||
@@ -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,12 +40,12 @@ 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
|
||||
profile_owned:
|
||||
ltx2_vae_tiling: generator.pipeline.profile_overrides.ltx2.vae_tiling
|
||||
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_size_in_pixels
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
|
||||
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_size_in_frames
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_overlap_in_frames
|
||||
ltx2_vae_tiling: generator.pipeline.vae_tiling
|
||||
preset_owned:
|
||||
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
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.preset_overrides.ltx2.vae.temporal_tile_overlap_in_frames
|
||||
ltx2_initial_latent_path: request.extensions.ltx2.initial_latent_path
|
||||
compatibility_only:
|
||||
mode: "Legacy multi-mode FastVideoArgs switch; typed inference config should not expose execution mode."
|
||||
@@ -69,16 +69,16 @@ surfaces:
|
||||
pipeline_config_base:
|
||||
moved:
|
||||
pipeline_config_path: generator.pipeline.components.pipeline_config_path
|
||||
profile_owned:
|
||||
embedded_cfg_scale: generator.pipeline.profile_overrides.embedded_cfg_scale
|
||||
flow_shift: generator.pipeline.profile_overrides.flow_shift
|
||||
flow_shift_sr: generator.pipeline.profile_overrides.flow_shift_sr
|
||||
is_causal: generator.pipeline.profile_overrides.is_causal
|
||||
vae_tiling: generator.pipeline.profile_overrides.vae_tiling
|
||||
vae_sp: generator.pipeline.profile_overrides.vae_sp
|
||||
dmd_denoising_steps: generator.pipeline.profile_overrides.dmd_denoising_steps
|
||||
ti2v_task: generator.pipeline.profile_overrides.ti2v_task
|
||||
boundary_ratio: generator.pipeline.profile_overrides.boundary_ratio
|
||||
preset_owned:
|
||||
embedded_cfg_scale: generator.pipeline.preset_overrides.embedded_cfg_scale
|
||||
flow_shift: generator.pipeline.preset_overrides.flow_shift
|
||||
flow_shift_sr: generator.pipeline.preset_overrides.flow_shift_sr
|
||||
is_causal: generator.pipeline.preset_overrides.is_causal
|
||||
vae_tiling: generator.pipeline.preset_overrides.vae_tiling
|
||||
vae_sp: generator.pipeline.preset_overrides.vae_sp
|
||||
dmd_denoising_steps: generator.pipeline.preset_overrides.dmd_denoising_steps
|
||||
ti2v_task: generator.pipeline.preset_overrides.ti2v_task
|
||||
boundary_ratio: generator.pipeline.preset_overrides.boundary_ratio
|
||||
compatibility_only:
|
||||
model_path: "Redundant with generator.model_path."
|
||||
disable_autocast: "Duplicated by generator.engine.disable_autocast during migration."
|
||||
@@ -97,7 +97,7 @@ surfaces:
|
||||
postprocess_text_funcs: "Internal text postprocessing hooks."
|
||||
|
||||
pipeline_config_extensions:
|
||||
profile_owned:
|
||||
preset_owned:
|
||||
conditioning_strategy:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.cosmos.CosmosConfig
|
||||
@@ -309,8 +309,8 @@ surfaces:
|
||||
compatibility_only:
|
||||
batch_size: "Gen3C inference-only tuning field pending typed batching design."
|
||||
gradient_checkpointing: "Gen3C inference-only compatibility field pending typed batching design."
|
||||
guidance_scale: "Gen3C pipeline-level default pending profile/default-request cleanup."
|
||||
num_inference_steps: "Gen3C pipeline-level default pending profile/default-request cleanup."
|
||||
guidance_scale: "Gen3C pipeline-level default pending preset/default-request cleanup."
|
||||
num_inference_steps: "Gen3C pipeline-level default pending preset/default-request cleanup."
|
||||
internal_only:
|
||||
audio_decoder_config: "Legacy internal component config object."
|
||||
audio_decoder_precision: "Precision override pending dedicated component precision design."
|
||||
@@ -345,6 +345,7 @@ surfaces:
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
num_inference_steps_sr: request.sampling.num_inference_steps_sr
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
@@ -353,96 +354,36 @@ surfaces:
|
||||
return_frames: request.output.return_frames
|
||||
return_trajectory_latents: request.runtime.return_trajectory_latents
|
||||
return_trajectory_decoded: request.runtime.return_trajectory_decoded
|
||||
profile_owned:
|
||||
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
|
||||
num_cond_frames: request.stage_overrides.refine.num_cond_frames
|
||||
trajectory_type: request.extensions.gen3c.trajectory_type
|
||||
movement_distance: request.extensions.gen3c.movement_distance
|
||||
camera_rotation: request.extensions.gen3c.camera_rotation
|
||||
prompt_attention_mask: request.extensions.hyworld.prompt_attention_mask
|
||||
negative_attention_mask: request.extensions.hyworld.negative_attention_mask
|
||||
camera_states: request.extensions.hunyuangamecraft.camera_states
|
||||
camera_trajectory: request.extensions.hunyuangamecraft.camera_trajectory
|
||||
action_list: request.extensions.hunyuangamecraft.action_list
|
||||
action_speed_list: request.extensions.hunyuangamecraft.action_speed_list
|
||||
gt_latents: request.extensions.hunyuangamecraft.gt_latents
|
||||
conditioning_mask: request.extensions.hunyuangamecraft.conditioning_mask
|
||||
ltx2_cfg_scale_video: request.extensions.ltx2.cfg_scale_video
|
||||
ltx2_cfg_scale_audio: request.extensions.ltx2.cfg_scale_audio
|
||||
ltx2_modality_scale_video: request.extensions.ltx2.modality_scale_video
|
||||
ltx2_modality_scale_audio: request.extensions.ltx2.modality_scale_audio
|
||||
ltx2_rescale_scale: request.extensions.ltx2.rescale_scale
|
||||
ltx2_stg_scale_video: request.extensions.ltx2.stg_scale_video
|
||||
ltx2_stg_scale_audio: request.extensions.ltx2.stg_scale_audio
|
||||
ltx2_stg_blocks_video: request.extensions.ltx2.stg_blocks_video
|
||||
ltx2_stg_blocks_audio: request.extensions.ltx2.stg_blocks_audio
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
|
||||
sampling_param_extensions:
|
||||
moved:
|
||||
guidance_scale_2:
|
||||
target: request.sampling.guidance_scale_2
|
||||
sources:
|
||||
- fastvideo.configs.sample.lingbotworld.LingBotWorld_SamplingParam
|
||||
- fastvideo.configs.sample.lingbotworld.Wan2_2_I2V_A14B_SamplingParam
|
||||
- fastvideo.configs.sample.wan.SelfForcingWan2_2_T2V_A14B_480P_SamplingParam
|
||||
- fastvideo.configs.sample.wan.Wan2_2_I2V_A14B_SamplingParam
|
||||
- fastvideo.configs.sample.wan.Wan2_2_T2V_A14B_SamplingParam
|
||||
profile_owned:
|
||||
action_list:
|
||||
target: request.extensions.hunyuangamecraft.action_list
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
action_speed_list:
|
||||
target: request.extensions.hunyuangamecraft.action_speed_list
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
camera_states:
|
||||
target: request.extensions.hunyuangamecraft.camera_states
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
camera_trajectory:
|
||||
target: request.extensions.hunyuangamecraft.camera_trajectory
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
conditioning_mask:
|
||||
target: request.extensions.hunyuangamecraft.conditioning_mask
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
gt_latents:
|
||||
target: request.extensions.hunyuangamecraft.gt_latents
|
||||
sources:
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
|
||||
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
|
||||
prompt_attention_mask:
|
||||
target: request.extensions.hyworld.prompt_attention_mask
|
||||
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
|
||||
negative_attention_mask:
|
||||
target: request.extensions.hyworld.negative_attention_mask
|
||||
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
|
||||
ltx2_cfg_scale_audio:
|
||||
target: request.extensions.ltx2.cfg_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_cfg_scale_video:
|
||||
target: request.extensions.ltx2.cfg_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_modality_scale_audio:
|
||||
target: request.extensions.ltx2.modality_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_modality_scale_video:
|
||||
target: request.extensions.ltx2.modality_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_rescale_scale:
|
||||
target: request.extensions.ltx2.rescale_scale
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_blocks_audio:
|
||||
target: request.extensions.ltx2.stg_blocks_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_blocks_video:
|
||||
target: request.extensions.ltx2.stg_blocks_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_scale_audio:
|
||||
target: request.extensions.ltx2.stg_scale_audio
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
ltx2_stg_scale_video:
|
||||
target: request.extensions.ltx2.stg_scale_video
|
||||
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
|
||||
sampling_param_extensions: {}
|
||||
|
||||
openai_image_request:
|
||||
kept:
|
||||
@@ -495,230 +436,14 @@ cli:
|
||||
notes:
|
||||
- "CLI parity is checked against the actual generate/serve parser dest sets."
|
||||
- "The inventory tracks parser dest names, excluding argparse's implicit help action."
|
||||
- "The refactored inference CLI is config-only: subcommands expose only --config, and any additional CLI input must use dotted override paths."
|
||||
generate:
|
||||
explicit_local_fields:
|
||||
- config
|
||||
expected_dests:
|
||||
- VSA_sparsity
|
||||
- boundary_ratio
|
||||
- bsa_cdf_threshold
|
||||
- bsa_chunk_k
|
||||
- bsa_chunk_q
|
||||
- bsa_sparsity
|
||||
- config
|
||||
- disable_autocast
|
||||
- dist_timeout
|
||||
- distributed_executor_backend
|
||||
- dit_config.prefix
|
||||
- dit_config.quant_config
|
||||
- dit_cpu_offload
|
||||
- dit_layerwise_offload
|
||||
- dit_precision
|
||||
- dmd_denoising_steps
|
||||
- embedded_cfg_scale
|
||||
- enable_bsa
|
||||
- enable_stage_verification
|
||||
- enable_torch_compile
|
||||
- flow_shift
|
||||
- fps
|
||||
- guidance_rescale
|
||||
- guidance_scale
|
||||
- height
|
||||
- hsdp_replicate_dim
|
||||
- hsdp_shard_dim
|
||||
- image_encoder_cpu_offload
|
||||
- image_encoder_precision
|
||||
- image_path
|
||||
- inference_mode
|
||||
- init_weights_from_safetensors
|
||||
- init_weights_from_safetensors_2
|
||||
- lora_nickname
|
||||
- lora_path
|
||||
- lora_target_modules
|
||||
- ltx2_initial_latent_path
|
||||
- ltx2_vae_spatial_tile_overlap_in_pixels
|
||||
- ltx2_vae_spatial_tile_size_in_pixels
|
||||
- ltx2_vae_temporal_tile_overlap_in_frames
|
||||
- ltx2_vae_temporal_tile_size_in_frames
|
||||
- ltx2_vae_tiling
|
||||
- master_port
|
||||
- moba_config_path
|
||||
- mode
|
||||
- model_path
|
||||
- negative_prompt
|
||||
- num_cond_frames
|
||||
- num_frames
|
||||
- num_gpus
|
||||
- num_inference_steps
|
||||
- num_videos_per_prompt
|
||||
- output_path
|
||||
- output_type
|
||||
- output_video_name
|
||||
- override_pipeline_cls_name
|
||||
- override_text_encoder_quant
|
||||
- override_text_encoder_safetensors
|
||||
- override_transformer_cls_name
|
||||
- pin_cpu_memory
|
||||
- pipeline_config_path
|
||||
- preprocess.dataloader_num_workers
|
||||
- preprocess.dataset_output_dir
|
||||
- preprocess.dataset_path
|
||||
- preprocess.dataset_type
|
||||
- preprocess.do_temporal_sample
|
||||
- preprocess.drop_short_ratio
|
||||
- preprocess.flush_frequency
|
||||
- preprocess.max_height
|
||||
- preprocess.max_width
|
||||
- preprocess.model_path
|
||||
- preprocess.num_frames
|
||||
- preprocess.preprocess_video_batch_size
|
||||
- preprocess.samples_per_file
|
||||
- preprocess.seed
|
||||
- preprocess.speed_factor
|
||||
- preprocess.train_fps
|
||||
- preprocess.training_cfg_rate
|
||||
- preprocess.video_length_tolerance_range
|
||||
- preprocess.video_loader_type
|
||||
- preprocess.with_audio
|
||||
- prompt
|
||||
- prompt_path
|
||||
- prompt_txt
|
||||
- refine_from
|
||||
- return_frames
|
||||
- return_trajectory_decoded
|
||||
- return_trajectory_latents
|
||||
- revision
|
||||
- save_video
|
||||
- seed
|
||||
- sp_size
|
||||
- spatial_refine_only
|
||||
- t_thresh
|
||||
- text_encoder_configs
|
||||
- text_encoder_cpu_offload
|
||||
- text_encoder_precisions
|
||||
- torch_compile_kwargs
|
||||
- tp_size
|
||||
- trust_remote_code
|
||||
- use_fsdp_inference
|
||||
- vae_config.blend_num_frames
|
||||
- vae_config.load_decoder
|
||||
- vae_config.load_encoder
|
||||
- vae_config.tile_sample_min_height
|
||||
- vae_config.tile_sample_min_num_frames
|
||||
- vae_config.tile_sample_min_width
|
||||
- vae_config.tile_sample_stride_height
|
||||
- vae_config.tile_sample_stride_num_frames
|
||||
- vae_config.tile_sample_stride_width
|
||||
- vae_config.use_parallel_tiling
|
||||
- vae_config.use_temporal_tiling
|
||||
- vae_config.use_tiling
|
||||
- vae_cpu_offload
|
||||
- vae_precision
|
||||
- vae_sp
|
||||
- vae_tiling
|
||||
- video_path
|
||||
- width
|
||||
- workload_type
|
||||
serve:
|
||||
explicit_local_fields:
|
||||
- config
|
||||
- host
|
||||
- output_dir
|
||||
- port
|
||||
expected_dests:
|
||||
- VSA_sparsity
|
||||
- bsa_cdf_threshold
|
||||
- bsa_chunk_k
|
||||
- bsa_chunk_q
|
||||
- bsa_sparsity
|
||||
- config
|
||||
- disable_autocast
|
||||
- dist_timeout
|
||||
- distributed_executor_backend
|
||||
- dit_config.prefix
|
||||
- dit_config.quant_config
|
||||
- dit_cpu_offload
|
||||
- dit_layerwise_offload
|
||||
- dit_precision
|
||||
- dmd_denoising_steps
|
||||
- embedded_cfg_scale
|
||||
- enable_bsa
|
||||
- enable_stage_verification
|
||||
- enable_torch_compile
|
||||
- flow_shift
|
||||
- host
|
||||
- hsdp_replicate_dim
|
||||
- hsdp_shard_dim
|
||||
- image_encoder_cpu_offload
|
||||
- image_encoder_precision
|
||||
- inference_mode
|
||||
- init_weights_from_safetensors
|
||||
- init_weights_from_safetensors_2
|
||||
- lora_nickname
|
||||
- lora_path
|
||||
- lora_target_modules
|
||||
- ltx2_initial_latent_path
|
||||
- ltx2_vae_spatial_tile_overlap_in_pixels
|
||||
- ltx2_vae_spatial_tile_size_in_pixels
|
||||
- ltx2_vae_temporal_tile_overlap_in_frames
|
||||
- ltx2_vae_temporal_tile_size_in_frames
|
||||
- ltx2_vae_tiling
|
||||
- master_port
|
||||
- mode
|
||||
- model_path
|
||||
- num_gpus
|
||||
- output_dir
|
||||
- output_type
|
||||
- override_pipeline_cls_name
|
||||
- override_text_encoder_quant
|
||||
- override_text_encoder_safetensors
|
||||
- override_transformer_cls_name
|
||||
- pin_cpu_memory
|
||||
- pipeline_config_path
|
||||
- port
|
||||
- preprocess.dataloader_num_workers
|
||||
- preprocess.dataset_output_dir
|
||||
- preprocess.dataset_path
|
||||
- preprocess.dataset_type
|
||||
- preprocess.do_temporal_sample
|
||||
- preprocess.drop_short_ratio
|
||||
- preprocess.flush_frequency
|
||||
- preprocess.max_height
|
||||
- preprocess.max_width
|
||||
- preprocess.model_path
|
||||
- preprocess.num_frames
|
||||
- preprocess.preprocess_video_batch_size
|
||||
- preprocess.samples_per_file
|
||||
- preprocess.seed
|
||||
- preprocess.speed_factor
|
||||
- preprocess.train_fps
|
||||
- preprocess.training_cfg_rate
|
||||
- preprocess.video_length_tolerance_range
|
||||
- preprocess.video_loader_type
|
||||
- preprocess.with_audio
|
||||
- prompt_txt
|
||||
- revision
|
||||
- sp_size
|
||||
- text_encoder_cpu_offload
|
||||
- text_encoder_precisions
|
||||
- torch_compile_kwargs
|
||||
- tp_size
|
||||
- trust_remote_code
|
||||
- use_fsdp_inference
|
||||
- vae_config.blend_num_frames
|
||||
- vae_config.load_decoder
|
||||
- vae_config.load_encoder
|
||||
- vae_config.tile_sample_min_height
|
||||
- vae_config.tile_sample_min_num_frames
|
||||
- vae_config.tile_sample_min_width
|
||||
- vae_config.tile_sample_stride_height
|
||||
- vae_config.tile_sample_stride_num_frames
|
||||
- vae_config.tile_sample_stride_width
|
||||
- vae_config.use_parallel_tiling
|
||||
- vae_config.use_temporal_tiling
|
||||
- vae_config.use_tiling
|
||||
- vae_cpu_offload
|
||||
- vae_precision
|
||||
- vae_sp
|
||||
- vae_tiling
|
||||
- workload_type
|
||||
|
||||
@@ -12,7 +12,7 @@ FastVideo maps a Diffusers-style repo into a pipeline like this:
|
||||
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
|
||||
weight name translation.
|
||||
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
|
||||
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
|
||||
- `fastvideo/api/sampling_param.py`: runtime sampling parameters.
|
||||
- `fastvideo/pipelines/basic/*`: end-to-end pipelines.
|
||||
- `fastvideo/pipelines/stages/*`: reusable pipeline stages.
|
||||
- `fastvideo/models/loader/*`: component loaders for Diffusers-style repos.
|
||||
@@ -26,7 +26,7 @@ Minimal usage (from `examples/inference/basic/basic.py`):
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
|
||||
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
|
||||
@@ -49,8 +49,9 @@ runtime parameters consistent:
|
||||
- `fastvideo/configs/models/`: architecture definitions, layer shapes, and
|
||||
`param_names_mapping` rules for key renaming.
|
||||
- `fastvideo/configs/pipelines/`: pipeline wiring and required components.
|
||||
- `fastvideo/configs/sample/`: default sampling parameters (steps, frames,
|
||||
guidance scale, resolution, fps).
|
||||
- `fastvideo/api/sampling_param.py`: sampling parameters (steps, frames,
|
||||
guidance scale, resolution, fps). Defaults come from profiles in
|
||||
`fastvideo/pipelines/basic/<family>/profiles.py`.
|
||||
- `fastvideo/registry.py`: unified registry for pipeline config + sampling
|
||||
defaults and model metadata resolution, defined via explicit
|
||||
`register_configs(...)` blocks (no separate dict registries).
|
||||
@@ -142,7 +143,7 @@ How this maps to FastVideo:
|
||||
- `T5TokenizerFast` -> loaded via HF in `fastvideo/models/loader/`
|
||||
- `UniPCMultistepScheduler` -> loaded via Diffusers scheduler utilities
|
||||
- Pipeline defaults -> `fastvideo/configs/pipelines/wan.py`
|
||||
- Sampling defaults -> `fastvideo/configs/sample/wan.py`
|
||||
- Sampling defaults -> `fastvideo/pipelines/basic/wan/profiles.py`
|
||||
|
||||
## Pipeline system
|
||||
|
||||
|
||||
@@ -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`.
|
||||
@@ -16,7 +16,8 @@ Both models are trained on **61×448×832** resolution but support generating vi
|
||||
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN \
|
||||
fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml
|
||||
```
|
||||
|
||||
## 🗂️ Dataset
|
||||
@@ -85,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.
|
||||
|
||||
@@ -33,7 +33,7 @@ The following two classes `PipelineConfig` and `SamplingParam` are used to confi
|
||||
|
||||
### SamplingParam
|
||||
|
||||
::: fastvideo.configs.sample.base.SamplingParam
|
||||
::: fastvideo.api.sampling_param.SamplingParam
|
||||
options:
|
||||
show_root_heading: true
|
||||
show_source: false
|
||||
|
||||
@@ -128,19 +128,14 @@ Concrete hierarchy: `DiTConfig` → `DiTArchConfig`, `VAEConfig` →
|
||||
- `dump_to_json()` / `load_from_json()` — JSON persistence. Callable
|
||||
fields and `arch_config` are excluded from dumps.
|
||||
|
||||
### SamplingParam (`fastvideo/configs/sample/`)
|
||||
### SamplingParam (`fastvideo/api/sampling_param.py`)
|
||||
|
||||
Generation parameters separate from pipeline config. Each model family
|
||||
provides defaults:
|
||||
provides defaults via a profile (see `fastvideo/pipelines/basic/<family>/profiles.py`):
|
||||
|
||||
```python
|
||||
@dataclass
|
||||
class WanT2V_1_3B_SamplingParam(SamplingParam):
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
guidance_scale: float = 3.0
|
||||
num_inference_steps: int = 50
|
||||
sp = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sp.height == 480, sp.width == 832, sp.num_frames == 81, etc.
|
||||
```
|
||||
|
||||
## Component Loading
|
||||
@@ -430,9 +425,9 @@ User: generator.generate_video(prompt, ...)
|
||||
`fastvideo/configs/pipelines/<model>.py`. Set DiT/VAE/encoder configs,
|
||||
flow_shift, precision defaults.
|
||||
|
||||
2. **Sampling param** — Create a `SamplingParam` subclass in
|
||||
`fastvideo/configs/sample/<model>.py`. Set default height, width,
|
||||
num_frames, guidance_scale, num_inference_steps.
|
||||
2. **Sampling param profile** — Create a profile in
|
||||
`fastvideo/pipelines/basic/<model>/profiles.py` with default height,
|
||||
width, num_frames, guidance_scale, num_inference_steps.
|
||||
|
||||
3. **Register configs** — In `fastvideo/registry.py`, add a
|
||||
`register_configs()` call inside `_register_configs()` with
|
||||
@@ -455,6 +450,6 @@ User: generator.generate_video(prompt, ...)
|
||||
`fastvideo/pipelines/stages/`, implement `forward()`, optionally
|
||||
implement `verify_input()`/`verify_output()`.
|
||||
|
||||
7. **Verify** — Run `fastvideo generate --model-path <path> --prompt
|
||||
"test" --num-inference-steps 2` to confirm the pipeline loads and
|
||||
generates output.
|
||||
7. **Verify** — Run `fastvideo generate --config <config.yaml>` with a
|
||||
minimal nested config to confirm the pipeline loads and generates
|
||||
output.
|
||||
|
||||
+42
-81
@@ -1,71 +1,29 @@
|
||||
# FastVideo CLI Inference
|
||||
|
||||
The FastVideo CLI exposes the same core inference controls as the Python API.
|
||||
The FastVideo CLI is config-first. Inference runs are driven by a nested JSON or
|
||||
YAML config, with optional dotted-path overrides on the command line. The
|
||||
contract matches training: use an explicit subcommand plus `--config`, then add
|
||||
any dotted overrides you need.
|
||||
|
||||
## Basic Usage
|
||||
|
||||
Use either:
|
||||
|
||||
1. `--model-path` + `--prompt`
|
||||
2. `--model-path` + `--prompt-txt` (batch prompts, one line per prompt)
|
||||
3. `--config` (JSON/YAML)
|
||||
|
||||
```bash
|
||||
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--prompt "A cat playing with a ball of yarn"
|
||||
fastvideo generate --config config.yaml
|
||||
fastvideo serve --config serve.yaml
|
||||
```
|
||||
|
||||
```bash
|
||||
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--prompt-txt prompts.txt
|
||||
```
|
||||
|
||||
You cannot provide both `--prompt` and `--prompt-txt` in the same run.
|
||||
|
||||
## View All Arguments
|
||||
|
||||
```bash
|
||||
fastvideo generate --help
|
||||
```
|
||||
|
||||
Arguments come from:
|
||||
The subcommands intentionally expose only `--config`. Any per-run CLI changes
|
||||
must use dotted override paths such as:
|
||||
|
||||
- FastVideo runtime args (`FastVideoArgs`)
|
||||
- Sampling args (`SamplingParam`)
|
||||
- Pipeline config args (`PipelineConfig`)
|
||||
|
||||
## Common Arguments
|
||||
|
||||
### Parallelism
|
||||
|
||||
- `--num-gpus`
|
||||
- `--sp-size`
|
||||
- `--tp-size`
|
||||
|
||||
### Sampling
|
||||
|
||||
- `--num-frames`
|
||||
- `--height` / `--width`
|
||||
- `--num-inference-steps`
|
||||
- `--guidance-scale`
|
||||
- `--seed`
|
||||
- `--negative-prompt`
|
||||
|
||||
### Output
|
||||
|
||||
- `--output-path`
|
||||
- `--save-video` / `--no-save-video`
|
||||
- `--return-frames`
|
||||
|
||||
### Offloading and Performance
|
||||
|
||||
- `--dit-layerwise-offload`
|
||||
- `--use-fsdp-inference`
|
||||
- `--text-encoder-cpu-offload`
|
||||
- `--image-encoder-cpu-offload`
|
||||
- `--vae-cpu-offload`
|
||||
- `--enable-torch-compile`
|
||||
- `--torch-compile-kwargs`
|
||||
- `--generator.engine.num_gpus 2`
|
||||
- `--request.sampling.seed 42`
|
||||
- `--server.port 9000`
|
||||
|
||||
## Using Config Files
|
||||
|
||||
@@ -73,50 +31,53 @@ Arguments come from:
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Config files can be JSON or YAML. CLI flags override config-file values.
|
||||
Config files can be JSON or YAML. Dotted CLI overrides take precedence over
|
||||
config-file values.
|
||||
|
||||
Example `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model_path: "FastVideo/FastHunyuan-diffusers"
|
||||
prompt: "A capybara lounging in a hammock"
|
||||
output_path: "outputs/"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
tp_size: 1
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
dit_precision: "bf16"
|
||||
vae_precision: "fp16"
|
||||
vae_tiling: true
|
||||
vae_sp: true
|
||||
enable_torch_compile: false
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
engine:
|
||||
num_gpus: 2
|
||||
parallelism:
|
||||
sp_size: 2
|
||||
tp_size: 1
|
||||
request:
|
||||
prompt: A capybara lounging in a hammock
|
||||
sampling:
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
output:
|
||||
output_path: outputs/
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- Use `dit_precision` / `vae_precision` (not `precision`).
|
||||
- Nested config objects are supported, for example `vae_config` and
|
||||
`dit_config`.
|
||||
- `generator` and `request` are the top-level keys for generation configs.
|
||||
- `serve` configs use `generator`, `server`, and optional `default_request`.
|
||||
- Prompt text files belong under `request.inputs.prompt_path`.
|
||||
|
||||
## Examples
|
||||
|
||||
Simple generation:
|
||||
|
||||
```bash
|
||||
fastvideo generate \
|
||||
--model-path FastVideo/FastHunyuan-diffusers \
|
||||
--prompt "A cat playing with a ball of yarn" \
|
||||
--num-frames 45 --height 720 --width 1280 \
|
||||
--num-inference-steps 6 --seed 1024 \
|
||||
--output-path outputs/
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Config + CLI override:
|
||||
Config + dotted override:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml --prompt "A panda skiing at sunset"
|
||||
fastvideo generate --config config.yaml --request.prompt "A panda skiing at sunset"
|
||||
```
|
||||
|
||||
Helper wrapper with positional config path:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/run.sh scripts/inference/inference_wan.yaml
|
||||
```
|
||||
|
||||
@@ -73,32 +73,40 @@ if __name__ == '__main__':
|
||||
|
||||
## JSON/YAML Config Files (CLI)
|
||||
|
||||
The CLI supports `--config` with JSON or YAML. Command-line arguments override
|
||||
config file values.
|
||||
By default, `fastvideo generate` uses `return_frames=false` unless you set
|
||||
`--return-frames` (or `return_frames: true` in config).
|
||||
The inference CLI is config-first. Use an explicit subcommand with `--config`,
|
||||
then apply optional dotted overrides on top, matching the training CLI style.
|
||||
By default, CLI generation uses `return_frames=false` unless you set
|
||||
`request.output.return_frames: true` in config or via a dotted override.
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml
|
||||
```
|
||||
|
||||
Use CLI argument names as keys (underscore or hyphen is accepted). Example:
|
||||
Example nested config:
|
||||
|
||||
```yaml
|
||||
model_path: "FastVideo/FastHunyuan-diffusers"
|
||||
prompt: "A capybara relaxing in a hammock"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
dit_precision: "bf16"
|
||||
vae_precision: "fp16"
|
||||
vae_tiling: true
|
||||
vae_sp: true
|
||||
enable_torch_compile: false
|
||||
generator:
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
engine:
|
||||
num_gpus: 2
|
||||
parallelism:
|
||||
sp_size: 2
|
||||
request:
|
||||
prompt: A capybara relaxing in a hammock
|
||||
sampling:
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
num_inference_steps: 6
|
||||
seed: 1024
|
||||
output:
|
||||
output_path: outputs/
|
||||
```
|
||||
|
||||
Override individual values from the CLI with dotted paths:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config config.yaml --request.sampling.seed 42
|
||||
```
|
||||
|
||||
## Performance Optimization
|
||||
|
||||
@@ -89,7 +89,7 @@ GEN3C defaults in FastVideo:
|
||||
|
||||
These values are defined in:
|
||||
|
||||
- `fastvideo/configs/sample/gen3c.py`
|
||||
- `fastvideo/pipelines/basic/gen3c/profiles.py`
|
||||
- `fastvideo/configs/pipelines/gen3c.py`
|
||||
|
||||
and align with the official GEN3C inference defaults in:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15_1080p"
|
||||
def main():
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
OUTPUT_PATH = "video_samples_lingbotworld"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
def main():
|
||||
|
||||
@@ -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()
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
|
||||
def main():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_1_Fun"
|
||||
OUTPUT_NAME = "wan2.1_test"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
|
||||
def main():
|
||||
|
||||
@@ -5,7 +5,7 @@ import time
|
||||
|
||||
import gradio as gr
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from copy import deepcopy
|
||||
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ import tempfile
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
MODEL_PATH_MAPPING = {
|
||||
|
||||
@@ -185,7 +185,7 @@ class BaseModelDeployment:
|
||||
|
||||
def _initialize_generator(self, config: Dict[str, Any]) -> None:
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
print(f"Initializing model: {self.model_path}")
|
||||
self.generator = VideoGenerator.from_pretrained(
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
Inference using a LoRA checkpoint from FastVideo trainer.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# Cosmos Predict2 2B T2V finetune config.
|
||||
#
|
||||
# Data must be preprocessed with Cosmos VAE + T5 text encoder
|
||||
# into parquet format before training.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: nvidia/Cosmos-Predict2-2B-Video2World
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos_preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
# Cosmos VAE: 4x temporal, 8x spatial compression.
|
||||
# 93 frames -> 24 latent frames, 480x832 -> 60x104
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 5000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos_finetune
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos
|
||||
run_name: cosmos_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos_pipeline.Cosmos2VideoToWorldPipeline
|
||||
dataset_file: data/cosmos_preprocessed/validation_prompts.json
|
||||
every_steps: 100
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -0,0 +1,79 @@
|
||||
# Cosmos-Predict2.5-2B Text-to-World overfitting test config.
|
||||
#
|
||||
# Overfits on a few short videos (480x832, 93 frames) to verify the
|
||||
# Cosmos 2.5 training plugin works end-to-end.
|
||||
#
|
||||
# Preprocess data first:
|
||||
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# bash examples/train/run.sh examples/train/configs/overfit_cosmos25_t2w.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: KyleShao/Cosmos-Predict2.5-2B-Diffusers
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
flow_shift: 1.0
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos25_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos25_overfit
|
||||
training_state_checkpointing_steps: 50
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos25
|
||||
run_name: cosmos25_overfit
|
||||
|
||||
model:
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos2_5_pipeline.Cosmos2_5Pipeline
|
||||
dataset_file: data/cosmos25_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 150
|
||||
sampling_steps: [35]
|
||||
guidance_scale: 7.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -10,6 +10,35 @@ set -ex
|
||||
|
||||
echo "Building fastvideo-kernel..."
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Neutralise conda-injected compiler toolchains.
|
||||
#
|
||||
# Conda compiler packages (gcc_linux-aarch64, gxx_linux-64, etc.) set
|
||||
# CMAKE_ARGS, CFLAGS, CXXFLAGS, and LDFLAGS on activation. When multiple
|
||||
# toolchains are installed the variables can reference a *cross*-compiler
|
||||
# that doesn't match the host (e.g. aarch64-conda-linux-gnu-c++ on x86_64).
|
||||
# Even when the correct toolchain is active, the flags it injects
|
||||
# (-march=nocona, -mtune=haswell, …) can conflict with nvcc's host-compiler
|
||||
# expectations. Clear them so CMake discovers the system compiler instead.
|
||||
# ---------------------------------------------------------------------------
|
||||
if [[ -n "${CONDA_PREFIX:-}" ]]; then
|
||||
_need_clean=0
|
||||
# Detect conda cross-compiler that doesn't match the host.
|
||||
_host_arch="$(uname -m)"
|
||||
if [[ "${CXX:-}" == *"conda"* ]] || [[ "${CC:-}" == *"conda"* ]]; then
|
||||
_need_clean=1
|
||||
fi
|
||||
if [[ "${CMAKE_ARGS:-}" == *"conda"* ]]; then
|
||||
_need_clean=1
|
||||
fi
|
||||
if (( _need_clean )); then
|
||||
echo "NOTE: Clearing conda-injected compiler settings (CC/CXX/CMAKE_ARGS/CFLAGS/...)"
|
||||
echo " to use the system compiler for CUDA extension builds."
|
||||
unset CC CXX CMAKE_ARGS CFLAGS CXXFLAGS LDFLAGS
|
||||
fi
|
||||
unset _need_clean _host_arch
|
||||
fi
|
||||
|
||||
# Ensure submodules are initialized if needed (tk)
|
||||
git submodule update --init --recursive
|
||||
|
||||
@@ -32,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
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
|
||||
@@ -7,21 +7,37 @@ from fastvideo.api.schema import (
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
GpuPoolConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
PlannedStage,
|
||||
PromptEnhancerConfig,
|
||||
PromptSafetyConfig,
|
||||
QuantizationConfig,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
ServerConfig,
|
||||
StreamingConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.presets import (
|
||||
InferencePreset,
|
||||
PresetStageSpec,
|
||||
get_all_preset_names,
|
||||
get_preset,
|
||||
get_presets_for_family,
|
||||
register_preset,
|
||||
validate_preset_selection,
|
||||
validate_stage_names,
|
||||
validate_stage_overrides,
|
||||
)
|
||||
from fastvideo.api.parser import (
|
||||
config_to_dict,
|
||||
load_config,
|
||||
@@ -31,6 +47,7 @@ from fastvideo.api.parser import (
|
||||
parse_config,
|
||||
)
|
||||
from fastvideo.api.results import GenerationResult
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
@@ -42,18 +59,26 @@ __all__ = [
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"GpuPoolConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"PromptEnhancerConfig",
|
||||
"PromptSafetyConfig",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"SamplingParam",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
"StreamingConfig",
|
||||
"WarmupConfig",
|
||||
"InferencePreset",
|
||||
"PresetStageSpec",
|
||||
"apply_overrides",
|
||||
"config_to_dict",
|
||||
"load_config",
|
||||
@@ -61,5 +86,12 @@ __all__ = [
|
||||
"load_run_config",
|
||||
"load_serve_config",
|
||||
"parse_cli_overrides",
|
||||
"get_all_preset_names",
|
||||
"get_preset",
|
||||
"get_presets_for_family",
|
||||
"parse_config",
|
||||
"register_preset",
|
||||
"validate_preset_selection",
|
||||
"validate_stage_names",
|
||||
"validate_stage_overrides",
|
||||
]
|
||||
|
||||
+216
-94
@@ -7,9 +7,17 @@ from dataclasses import fields, is_dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.overrides import apply_overrides, normalize_overrides
|
||||
from fastvideo.api.parser import config_to_dict, load_raw_config, parse_config
|
||||
from fastvideo.api.request_metadata import (
|
||||
EXPLICIT_PATHS_ATTR,
|
||||
bind_generation_request_raw,
|
||||
get_explicit_paths,
|
||||
reset_tracking_roots,
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
@@ -17,11 +25,14 @@ from fastvideo.api.schema import (
|
||||
RequestRuntimeConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
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
|
||||
|
||||
_EXPLICIT_REQUEST_ATTR = "_fastvideo_explicit_request"
|
||||
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
|
||||
_SAMPLING_FIELD_NAMES = {field.name for field in fields(SamplingConfig)}
|
||||
_RUNTIME_FIELD_NAMES = {field.name for field in fields(RequestRuntimeConfig)}
|
||||
@@ -33,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:
|
||||
@@ -46,7 +61,7 @@ def load_generator_config_from_file(
|
||||
overrides: list[str] | Mapping[str, Any] | None = None,
|
||||
) -> GeneratorConfig:
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = _normalize_overrides(overrides)
|
||||
normalized_overrides = normalize_overrides(overrides)
|
||||
|
||||
if _looks_like_run_or_serve_config(raw):
|
||||
if normalized_overrides:
|
||||
@@ -75,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":
|
||||
@@ -101,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":
|
||||
@@ -142,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:
|
||||
@@ -153,16 +199,12 @@ def legacy_from_pretrained_to_config(
|
||||
def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, Any], ) -> FastVideoArgs:
|
||||
normalized = normalize_generator_config(config)
|
||||
unsupported = []
|
||||
if normalized.pipeline.profile is not None:
|
||||
unsupported.append("pipeline.profile")
|
||||
if normalized.pipeline.profile_version is not None:
|
||||
unsupported.append("pipeline.profile_version")
|
||||
if normalized.pipeline.components.config_root is not None:
|
||||
unsupported.append("pipeline.components.config_root")
|
||||
if normalized.pipeline.preset is not None:
|
||||
unsupported.append("pipeline.preset")
|
||||
if normalized.pipeline.preset_version is not None:
|
||||
unsupported.append("pipeline.preset_version")
|
||||
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")
|
||||
@@ -186,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:
|
||||
@@ -215,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.profile_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)
|
||||
|
||||
@@ -224,8 +284,10 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
def normalize_generation_request(request: GenerationRequest | Mapping[str, Any], ) -> GenerationRequest:
|
||||
normalized = (request if isinstance(request, GenerationRequest) else parse_config(GenerationRequest, request))
|
||||
|
||||
if not hasattr(normalized, _EXPLICIT_REQUEST_ATTR):
|
||||
setattr(normalized, _EXPLICIT_REQUEST_ATTR, _serialize_generation_request(normalized))
|
||||
if not hasattr(normalized, EXPLICIT_PATHS_ATTR):
|
||||
# Request wasn't bound through the parser (e.g. constructed
|
||||
# directly). Treat every currently-set field as explicit.
|
||||
bind_generation_request_raw(normalized, _serialize_generation_request(normalized))
|
||||
return normalized
|
||||
|
||||
|
||||
@@ -253,7 +315,7 @@ def legacy_generate_call_to_request(
|
||||
raw.setdefault("inputs", {})["grid_sizes"] = grid_sizes
|
||||
|
||||
normalized = parse_config(GenerationRequest, raw)
|
||||
setattr(normalized, _EXPLICIT_REQUEST_ATTR, deepcopy(raw))
|
||||
bind_generation_request_raw(normalized, raw)
|
||||
return normalized
|
||||
|
||||
|
||||
@@ -264,16 +326,24 @@ 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)
|
||||
updates = _explicit_request_updates(request)
|
||||
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():
|
||||
if hasattr(sampling_param, key):
|
||||
setattr(sampling_param, key, deepcopy(value))
|
||||
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS or _is_supported_as_default_only(key, value):
|
||||
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
|
||||
continue
|
||||
elif value == _SCHEMA_DEFAULT_UPDATES.get(key, _MISSING):
|
||||
# Schema-default field that isn't on SamplingParam; tolerated
|
||||
# because direct GenerationRequest(...) construction has no
|
||||
# way to distinguish "user set" from "schema default".
|
||||
continue
|
||||
else:
|
||||
raise ValueError(f"Request field {key!r} is not supported by sampling params for {model_path}")
|
||||
@@ -290,10 +360,12 @@ def expand_request_prompt_batch(request: GenerationRequest, ) -> list[Generation
|
||||
requests: list[GenerationRequest] = []
|
||||
for index, prompt in enumerate(request.prompt):
|
||||
single_request = deepcopy(request)
|
||||
# deepcopy preserves the tracking-root cycle, but re-pin roots
|
||||
# defensively so that subsequent setattrs record on the copy.
|
||||
reset_tracking_roots(single_request)
|
||||
single_request.prompt = prompt
|
||||
_fan_out_batched_input_value(request, single_request, "image_path", index)
|
||||
_fan_out_batched_input_value(request, single_request, "video_path", index)
|
||||
_fan_out_explicit_request_metadata(request, single_request, index, prompt)
|
||||
requests.append(single_request)
|
||||
return requests
|
||||
|
||||
@@ -302,12 +374,23 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
|
||||
return isinstance(raw.get("generator"), Mapping)
|
||||
|
||||
|
||||
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
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]:
|
||||
@@ -348,20 +431,83 @@ def _apply_request_field(
|
||||
|
||||
def request_to_pipeline_overrides(request: GenerationRequest) -> dict[str, Any]:
|
||||
overrides: dict[str, Any] = {}
|
||||
for key, value in _explicit_request_updates(request).items():
|
||||
for key, value in explicit_request_updates(request).items():
|
||||
if key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
|
||||
overrides[key] = deepcopy(value)
|
||||
return overrides
|
||||
|
||||
|
||||
def _explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
|
||||
raw = getattr(request, _EXPLICIT_REQUEST_ATTR, None)
|
||||
if raw is None:
|
||||
raw = _serialize_generation_request(request)
|
||||
def explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
|
||||
"""Project a ``GenerationRequest`` down to *explicitly set* fields only.
|
||||
|
||||
Returns a flat kwargs dict suitable for merging into a generator call.
|
||||
The projection uses ``_fastvideo_explicit_paths`` (populated during
|
||||
``parse_config`` / raw binding) so schema defaults on the dataclass
|
||||
are **not** emitted — only paths the caller/operator actually wrote.
|
||||
|
||||
This is what makes ``ServeConfig.default_request`` work as an
|
||||
operator-pinned baseline rather than a full override: a YAML with just
|
||||
``sampling.seed: 42`` yields ``{"seed": 42}``, not the full sampling
|
||||
config with its 15 schema defaults.
|
||||
|
||||
Precondition: the request must carry ``_fastvideo_explicit_paths`` —
|
||||
populated by :func:`fastvideo.api.parser.parse_config` or
|
||||
:func:`fastvideo.api.compat.normalize_generation_request`. Calling on
|
||||
a raw ``GenerationRequest()`` asserts.
|
||||
"""
|
||||
assert hasattr(request,
|
||||
EXPLICIT_PATHS_ATTR), ("GenerationRequest reached explicit_request_updates without tracking; "
|
||||
"every entry point must route through normalize_generation_request "
|
||||
"or parse_config first")
|
||||
paths = get_explicit_paths(request)
|
||||
raw = _build_sparse_raw_from_paths(request, paths)
|
||||
return _extract_request_updates(raw)
|
||||
|
||||
|
||||
def _build_sparse_raw_from_paths(
|
||||
request: GenerationRequest,
|
||||
paths: frozenset[str],
|
||||
) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {}
|
||||
for path in paths:
|
||||
parts = path.split(".")
|
||||
value = _read_dotted_path(request, parts)
|
||||
if value is _MISSING:
|
||||
continue
|
||||
_set_dotted_path(result, parts, deepcopy(value))
|
||||
return result
|
||||
|
||||
|
||||
def _read_dotted_path(obj: Any, parts: list[str]) -> Any:
|
||||
for part in parts:
|
||||
if is_dataclass(obj) and not isinstance(obj, type):
|
||||
if not hasattr(obj, part):
|
||||
return _MISSING
|
||||
obj = getattr(obj, part)
|
||||
elif isinstance(obj, Mapping):
|
||||
if part not in obj:
|
||||
return _MISSING
|
||||
obj = obj[part]
|
||||
else:
|
||||
return _MISSING
|
||||
return obj
|
||||
|
||||
|
||||
def _set_dotted_path(
|
||||
target: dict[str, Any],
|
||||
parts: list[str],
|
||||
value: Any,
|
||||
) -> None:
|
||||
cursor = target
|
||||
for part in parts[:-1]:
|
||||
nxt = cursor.get(part)
|
||||
if not isinstance(nxt, dict):
|
||||
nxt = {}
|
||||
cursor[part] = nxt
|
||||
cursor = nxt
|
||||
cursor[parts[-1]] = value
|
||||
|
||||
|
||||
def _extract_request_updates(raw: Mapping[str, Any]) -> dict[str, Any]:
|
||||
updates: dict[str, Any] = {}
|
||||
if "negative_prompt" in raw:
|
||||
@@ -405,6 +551,40 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
|
||||
return deepcopy(config_to_dict(request))
|
||||
|
||||
|
||||
_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,
|
||||
target_request: GenerationRequest,
|
||||
@@ -418,29 +598,6 @@ def _fan_out_batched_input_value(
|
||||
setattr(target_request.inputs, field_name, deepcopy(value[index]))
|
||||
|
||||
|
||||
def _fan_out_explicit_request_metadata(
|
||||
source_request: GenerationRequest,
|
||||
target_request: GenerationRequest,
|
||||
index: int,
|
||||
prompt: str,
|
||||
) -> None:
|
||||
raw = getattr(source_request, _EXPLICIT_REQUEST_ATTR, None)
|
||||
if raw is None:
|
||||
return
|
||||
|
||||
raw = deepcopy(raw)
|
||||
raw["prompt"] = prompt
|
||||
inputs = raw.get("inputs")
|
||||
if isinstance(inputs, dict):
|
||||
for field_name in ("image_path", "video_path"):
|
||||
value = inputs.get(field_name)
|
||||
if isinstance(value, list):
|
||||
_validate_batched_input_length(source_request.prompt, value, field_name)
|
||||
inputs[field_name] = deepcopy(value[index])
|
||||
|
||||
setattr(target_request, _EXPLICIT_REQUEST_ATTR, raw)
|
||||
|
||||
|
||||
def _validate_batched_input_length(
|
||||
prompts: str | list[str] | None,
|
||||
values: list[Any],
|
||||
@@ -452,50 +609,15 @@ def _validate_batched_input_length(
|
||||
raise ValueError(f"GenerationRequest.inputs.{field_name} must have the same length as request.prompt")
|
||||
|
||||
|
||||
def _is_supported_as_default_only(key: str, value: Any) -> bool:
|
||||
default_value = _DEFAULT_REQUEST_UPDATES.get(key, _MISSING)
|
||||
return default_value is not _MISSING and _values_equal(value, default_value)
|
||||
|
||||
|
||||
def _collect_non_default_fields(
|
||||
value: Any,
|
||||
default: Any,
|
||||
) -> dict[str, Any]:
|
||||
if not (is_dataclass(value) and is_dataclass(default)):
|
||||
return {}
|
||||
|
||||
result: dict[str, Any] = {}
|
||||
for field in fields(value):
|
||||
current = getattr(value, field.name)
|
||||
default_value = getattr(default, field.name)
|
||||
if is_dataclass(current) and is_dataclass(default_value):
|
||||
nested = _collect_non_default_fields(current, default_value)
|
||||
if nested:
|
||||
result[field.name] = nested
|
||||
continue
|
||||
if not _values_equal(current, default_value):
|
||||
result[field.name] = deepcopy(current)
|
||||
return result
|
||||
|
||||
|
||||
def _values_equal(left: Any, right: Any) -> bool:
|
||||
if left is right:
|
||||
return True
|
||||
try:
|
||||
return bool(left == right)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
_DEFAULT_REQUEST_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
|
||||
|
||||
__all__ = [
|
||||
"explicit_request_updates",
|
||||
"generator_config_to_fastvideo_args",
|
||||
"legacy_from_pretrained_to_config",
|
||||
"legacy_generate_call_to_request",
|
||||
"load_generator_config_from_file",
|
||||
"normalize_generation_request",
|
||||
"normalize_generator_config",
|
||||
"register_continuation_kind",
|
||||
"request_to_pipeline_overrides",
|
||||
"request_to_sampling_param",
|
||||
]
|
||||
|
||||
@@ -31,7 +31,7 @@ def parse_cli_overrides(overrides: list[str]) -> dict[str, Any]:
|
||||
raise ValueError(f"Missing value for override {token!r}")
|
||||
raw_value = overrides[index]
|
||||
|
||||
parsed[key] = _cast_override_value(raw_value)
|
||||
parsed[_normalize_override_key(key)] = _cast_override_value(raw_value)
|
||||
index += 1
|
||||
|
||||
return parsed
|
||||
@@ -45,6 +45,15 @@ def apply_overrides(config: Mapping[str, Any], overrides: Mapping[str, Any]) ->
|
||||
return merged
|
||||
|
||||
|
||||
def normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
"""Normalize a CLI list or mapping of overrides into a flat dict."""
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
|
||||
|
||||
def _apply_single_override(config: dict[str, Any], dotted_key: str, value: Any) -> None:
|
||||
parts = dotted_key.split(".")
|
||||
if not all(parts):
|
||||
@@ -94,4 +103,8 @@ def _cast_override_value(raw: str) -> Any:
|
||||
return raw
|
||||
|
||||
|
||||
__all__ = ["apply_overrides", "parse_cli_overrides"]
|
||||
def _normalize_override_key(key: str) -> str:
|
||||
return key.replace("-", "_")
|
||||
|
||||
|
||||
__all__ = ["apply_overrides", "normalize_overrides", "parse_cli_overrides"]
|
||||
|
||||
+16
-12
@@ -11,8 +11,13 @@ from typing import Any, Literal, TypeVar, Union, get_args, get_origin, get_type_
|
||||
import yaml
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.schema import RunConfig, ServeConfig
|
||||
from fastvideo.api.overrides import apply_overrides, normalize_overrides
|
||||
from fastvideo.api.request_metadata import (
|
||||
bind_generation_request_raw,
|
||||
bind_run_config_raw,
|
||||
bind_serve_config_raw,
|
||||
)
|
||||
from fastvideo.api.schema import GenerationRequest, RunConfig, ServeConfig
|
||||
|
||||
T = TypeVar("T")
|
||||
_UNION_ORIGINS = {types.UnionType, Union}
|
||||
@@ -31,7 +36,14 @@ def parse_config(config_type: type[T], raw: Mapping[str, Any] | T) -> T:
|
||||
return raw
|
||||
if not isinstance(raw, Mapping):
|
||||
raise ConfigValidationError("", f"expected mapping for {config_type.__name__}")
|
||||
return _SchemaParser().parse_dataclass(config_type, raw, "")
|
||||
parsed = _SchemaParser().parse_dataclass(config_type, raw, "")
|
||||
if config_type is GenerationRequest:
|
||||
return bind_generation_request_raw(parsed, raw)
|
||||
if config_type is RunConfig:
|
||||
return bind_run_config_raw(parsed, raw)
|
||||
if config_type is ServeConfig:
|
||||
return bind_serve_config_raw(parsed, raw)
|
||||
return parsed
|
||||
|
||||
|
||||
def config_to_dict(config: Any) -> Any:
|
||||
@@ -52,7 +64,7 @@ def load_config(
|
||||
) -> T:
|
||||
"""Load a typed config object from YAML or JSON."""
|
||||
raw = load_raw_config(path)
|
||||
normalized_overrides = _normalize_overrides(overrides)
|
||||
normalized_overrides = normalize_overrides(overrides)
|
||||
if normalized_overrides:
|
||||
raw = apply_overrides(raw, normalized_overrides)
|
||||
return parse_config(config_type, raw)
|
||||
@@ -96,14 +108,6 @@ def _load_raw_mapping(handle: Any, config_path: Path) -> Any:
|
||||
raise ValueError(f"Unsupported config file format: {config_path}")
|
||||
|
||||
|
||||
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
|
||||
if not overrides:
|
||||
return None
|
||||
if isinstance(overrides, list):
|
||||
return parse_cli_overrides(overrides)
|
||||
return dict(overrides)
|
||||
|
||||
|
||||
class _SchemaParser:
|
||||
|
||||
def parse_dataclass(
|
||||
|
||||
@@ -0,0 +1,261 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pipeline preset registry.
|
||||
|
||||
A *preset* is a named inference preset for a model family. It bundles:
|
||||
|
||||
* ``defaults`` — sampling values applied when the user does not
|
||||
override them (consumed at runtime via ``SamplingParam.from_pretrained``);
|
||||
* ``stage_schemas`` — **validation-only** metadata describing which
|
||||
user-facing stage names (``"denoise"``, ``"sr"``) the preset recognises
|
||||
and which ``stage_overrides`` keys each stage accepts.
|
||||
|
||||
The ``stage_schemas`` tuple does **not** drive pipeline execution. The
|
||||
concrete execution DAG (text encoding, denoising, VAE decoding, …) is
|
||||
hard-coded per-pipeline in ``create_pipeline_stages()``. Schemas exist
|
||||
purely so that ``PipelineSelection.preset`` and
|
||||
``GenerationRequest.stage_overrides`` can be type-checked up front
|
||||
without touching the pipeline.
|
||||
|
||||
Preset base types and the registry API live here (public API surface).
|
||||
Preset *instances* are defined in pipeline-local ``presets.py`` files
|
||||
(e.g. ``fastvideo/pipelines/basic/wan/presets.py``) and registered
|
||||
explicitly from :func:`_register_presets` in ``fastvideo/registry.py``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.errors import ConfigValidationError
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Types
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PresetStageSpec:
|
||||
"""A user-facing stage name within a preset, used only to validate
|
||||
``stage_overrides`` keys. Not read by pipeline execution — the real
|
||||
execution DAG lives in each pipeline's ``create_pipeline_stages()``.
|
||||
"""
|
||||
|
||||
name: str
|
||||
"""Short user-facing name, e.g. ``"denoise"``, ``"sr"``."""
|
||||
|
||||
kind: str
|
||||
"""Semantic kind, e.g. ``"denoising"``, ``"super_resolution"``."""
|
||||
|
||||
description: str = ""
|
||||
|
||||
allowed_overrides: frozenset[str] = field(default_factory=frozenset)
|
||||
"""Keys that may appear in ``stage_overrides[name]``."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class InferencePreset:
|
||||
"""A named inference preset for a model family."""
|
||||
|
||||
name: str
|
||||
"""Preset name, e.g. ``"wan_t2v_1_3b"``."""
|
||||
|
||||
version: int
|
||||
"""Preset schema version; bump on breaking schema changes."""
|
||||
|
||||
model_family: str
|
||||
"""Model family key, e.g. ``"wan"``, ``"ltx2"``."""
|
||||
|
||||
description: str = ""
|
||||
|
||||
workload_type: str | None = None
|
||||
"""Optional workload hint: ``"t2v"``, ``"i2v"``, etc."""
|
||||
|
||||
stage_schemas: tuple[PresetStageSpec, ...] = ()
|
||||
"""User-facing stage names for ``stage_overrides`` validation.
|
||||
|
||||
Validation-only: this tuple is consumed by
|
||||
:func:`validate_stage_overrides` and is **not** used to drive
|
||||
pipeline execution. Omit or leave empty if the preset exposes no
|
||||
per-stage override surface.
|
||||
"""
|
||||
|
||||
defaults: dict[str, Any] = field(default_factory=dict)
|
||||
"""Preset-level default sampling/runtime values."""
|
||||
|
||||
stage_defaults: dict[str, dict[str, Any]] = field(default_factory=dict)
|
||||
"""Per-stage default overrides, keyed by stage name."""
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Registry
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
# Keyed by (model_family, name, version).
|
||||
_PRESET_REGISTRY: dict[tuple[str, str, int], InferencePreset] = {}
|
||||
|
||||
|
||||
def register_preset(preset: InferencePreset) -> None:
|
||||
"""Register a preset definition.
|
||||
|
||||
Raises :class:`ValueError` on duplicate
|
||||
``(model_family, name, version)`` keys.
|
||||
"""
|
||||
key = (preset.model_family, preset.name, preset.version)
|
||||
if key in _PRESET_REGISTRY:
|
||||
raise ValueError(f"Duplicate preset registration: "
|
||||
f"model_family={key[0]!r}, name={key[1]!r}, "
|
||||
f"version={key[2]!r}")
|
||||
_PRESET_REGISTRY[key] = preset
|
||||
|
||||
|
||||
def get_preset(
|
||||
name: str,
|
||||
model_family: str,
|
||||
version: int | None = None,
|
||||
) -> InferencePreset:
|
||||
"""Look up a registered preset.
|
||||
|
||||
When *version* is ``None`` the highest registered version for the
|
||||
given *(model_family, name)* pair is returned.
|
||||
|
||||
Raises :class:`~fastvideo.api.errors.ConfigValidationError` when the
|
||||
preset cannot be found.
|
||||
"""
|
||||
if version is not None:
|
||||
key = (model_family, name, version)
|
||||
preset = _PRESET_REGISTRY.get(key)
|
||||
if preset is not None:
|
||||
return preset
|
||||
raise ConfigValidationError(
|
||||
"pipeline.preset",
|
||||
f"unknown preset {name!r} version {version!r} "
|
||||
f"for model family {model_family!r}; "
|
||||
f"registered: {_format_registered(model_family)}",
|
||||
)
|
||||
|
||||
# Find the highest version for (model_family, name).
|
||||
candidates = [prof for (fam, n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family and n == name]
|
||||
if not candidates:
|
||||
raise ConfigValidationError(
|
||||
"pipeline.preset",
|
||||
f"unknown preset {name!r} for model family "
|
||||
f"{model_family!r}; "
|
||||
f"registered: {_format_registered(model_family)}",
|
||||
)
|
||||
return max(candidates, key=lambda p: p.version)
|
||||
|
||||
|
||||
def get_presets_for_family(model_family: str, ) -> list[InferencePreset]:
|
||||
"""Return all presets registered for *model_family*."""
|
||||
return [prof for (fam, _n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family]
|
||||
|
||||
|
||||
def get_all_preset_names() -> list[str]:
|
||||
"""Return the sorted list of all registered preset names."""
|
||||
return sorted({prof.name for prof in _PRESET_REGISTRY.values()})
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Validation helpers
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
def validate_stage_names(
|
||||
preset: InferencePreset,
|
||||
stage_overrides: Mapping[str, Any],
|
||||
) -> None:
|
||||
"""Check that *stage_overrides* keys are valid stage names.
|
||||
|
||||
Raises :class:`~fastvideo.api.errors.ConfigValidationError` with a
|
||||
path-qualified message for unknown stage names.
|
||||
"""
|
||||
valid_names = {stage.name for stage in preset.stage_schemas}
|
||||
for stage_name in stage_overrides:
|
||||
if stage_name not in valid_names:
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}",
|
||||
f"unknown stage for preset {preset.name!r}; "
|
||||
f"valid stages: {sorted(valid_names)}",
|
||||
)
|
||||
|
||||
|
||||
def validate_stage_overrides(
|
||||
preset: InferencePreset,
|
||||
stage_overrides: Mapping[str, Any],
|
||||
) -> None:
|
||||
"""Validate stage override keys against the preset.
|
||||
|
||||
Calls :func:`validate_stage_names` first, then checks that each
|
||||
override key is in the stage's ``allowed_overrides``.
|
||||
"""
|
||||
validate_stage_names(preset, stage_overrides)
|
||||
stages_by_name = {stage.name: stage for stage in preset.stage_schemas}
|
||||
for stage_name, overrides in stage_overrides.items():
|
||||
if not isinstance(overrides, Mapping):
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}",
|
||||
"must be a mapping",
|
||||
)
|
||||
stage_spec = stages_by_name[stage_name]
|
||||
if not stage_spec.allowed_overrides:
|
||||
if overrides:
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}",
|
||||
f"stage {stage_name!r} does not accept "
|
||||
f"overrides",
|
||||
)
|
||||
continue
|
||||
for key in overrides:
|
||||
if key not in stage_spec.allowed_overrides:
|
||||
raise ConfigValidationError(
|
||||
f"stage_overrides.{stage_name}.{key}",
|
||||
f"not an allowed override for stage "
|
||||
f"{stage_name!r}; allowed: "
|
||||
f"{sorted(stage_spec.allowed_overrides)}",
|
||||
)
|
||||
|
||||
|
||||
def validate_preset_selection(
|
||||
preset_name: str | None,
|
||||
model_family: str,
|
||||
*,
|
||||
preset_version: int | None = None,
|
||||
stage_overrides: Mapping[str, Any] | None = None,
|
||||
) -> InferencePreset | None:
|
||||
"""Resolve and validate a preset selection end-to-end.
|
||||
|
||||
Returns the resolved :class:`InferencePreset`, or ``None`` if
|
||||
*preset_name* is ``None`` (no preset requested).
|
||||
"""
|
||||
if preset_name is None:
|
||||
return None
|
||||
preset = get_preset(preset_name, model_family, version=preset_version)
|
||||
if stage_overrides:
|
||||
validate_stage_overrides(preset, stage_overrides)
|
||||
return preset
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
def _format_registered(model_family: str) -> str:
|
||||
names = sorted({prof.name for (fam, _n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family})
|
||||
if not names:
|
||||
return "(none)"
|
||||
return ", ".join(repr(n) for n in names)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InferencePreset",
|
||||
"PresetStageSpec",
|
||||
"get_all_preset_names",
|
||||
"get_preset",
|
||||
"get_presets_for_family",
|
||||
"register_preset",
|
||||
"validate_preset_selection",
|
||||
"validate_stage_names",
|
||||
"validate_stage_overrides",
|
||||
]
|
||||
@@ -0,0 +1,233 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Track which GenerationRequest fields the user explicitly provided.
|
||||
|
||||
When translating a GenerationRequest into a legacy SamplingParam we must
|
||||
distinguish user-provided values (which should override model defaults)
|
||||
from schema defaults (which should NOT override model defaults).
|
||||
|
||||
The mechanism: a single ``_fastvideo_explicit_paths`` set stored on the
|
||||
root ``GenerationRequest``. It holds dotted leaf paths (e.g.
|
||||
``"sampling.guidance_scale"``) the user has touched, either via raw
|
||||
config at bind time or via attribute assignment at runtime. A patched
|
||||
``__setattr__`` on the request dataclass types records assignments into
|
||||
this set.
|
||||
|
||||
The set holds leaf paths only. Nested dataclass or mapping assignments
|
||||
are flattened to their leaves at record time.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Mapping
|
||||
import dataclasses
|
||||
from typing import Any, cast
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
PlannedStage,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
|
||||
EXPLICIT_PATHS_ATTR = "_fastvideo_explicit_paths"
|
||||
|
||||
_TRACKING_ROOT_ATTR = "_fastvideo_request_tracking_root"
|
||||
_TRACKING_PATH_ATTR = "_fastvideo_request_tracking_path"
|
||||
_TRACKING_PATCHED_ATTR = "_fastvideo_request_tracking_patched"
|
||||
|
||||
_TRACKED_REQUEST_TYPES = (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
SamplingConfig,
|
||||
RequestRuntimeConfig,
|
||||
OutputConfig,
|
||||
ContinuationState,
|
||||
PlannedStage,
|
||||
GenerationPlan,
|
||||
)
|
||||
|
||||
|
||||
def bind_generation_request_raw(
|
||||
request: GenerationRequest,
|
||||
raw: Mapping[str, Any] | None,
|
||||
) -> GenerationRequest:
|
||||
"""Install explicit-path tracking on *request*.
|
||||
|
||||
*raw* is the parsed config dict (YAML/JSON/kwargs); every leaf key
|
||||
in it becomes an explicit path. Subsequent attribute assignments on
|
||||
*request* or its nested dataclasses are recorded automatically via a
|
||||
patched ``__setattr__``.
|
||||
"""
|
||||
_ensure_request_tracking()
|
||||
# Disable recording while we walk the tree to install roots.
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, None)
|
||||
_set_tracking_roots(request, request, "")
|
||||
paths: set[str] = set()
|
||||
_record_value_paths(raw or {}, "", paths)
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, paths)
|
||||
return request
|
||||
|
||||
|
||||
def bind_run_config_raw(
|
||||
config: RunConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> RunConfig:
|
||||
request_raw = raw.get("request")
|
||||
if isinstance(request_raw, Mapping):
|
||||
bind_generation_request_raw(config.request, request_raw)
|
||||
else:
|
||||
bind_generation_request_raw(config.request, {})
|
||||
return config
|
||||
|
||||
|
||||
def bind_serve_config_raw(
|
||||
config: ServeConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> ServeConfig:
|
||||
default_request_raw = raw.get("default_request")
|
||||
if isinstance(default_request_raw, Mapping):
|
||||
bind_generation_request_raw(config.default_request, default_request_raw)
|
||||
else:
|
||||
bind_generation_request_raw(config.default_request, {})
|
||||
return config
|
||||
|
||||
|
||||
def get_explicit_paths(request: GenerationRequest) -> frozenset[str]:
|
||||
"""Return a snapshot of the explicit paths set on *request*."""
|
||||
paths = getattr(request, EXPLICIT_PATHS_ATTR, None)
|
||||
if isinstance(paths, set | frozenset):
|
||||
return frozenset(paths)
|
||||
return frozenset()
|
||||
|
||||
|
||||
def reset_tracking_roots(request: GenerationRequest) -> None:
|
||||
"""Re-install tracking roots after a deepcopy or manual clone.
|
||||
|
||||
The paths set itself deepcopies correctly; we only need to repoint
|
||||
the tracking root on nested dataclasses at the new root.
|
||||
"""
|
||||
_ensure_request_tracking()
|
||||
_set_tracking_roots(request, request, "")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Path recording
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _record_value_paths(
|
||||
value: Any,
|
||||
prefix: str,
|
||||
out: set[str],
|
||||
) -> None:
|
||||
"""Add every leaf path under *value* to *out*.
|
||||
|
||||
A leaf is any terminal value (non-dataclass, non-mapping, or empty
|
||||
mapping/dataclass). ``prefix`` is the dotted path at which *value*
|
||||
sits. When called with an empty ``prefix`` (the root), leaves are
|
||||
recorded at their own key.
|
||||
"""
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
dc_fields = dataclasses.fields(value)
|
||||
if not dc_fields:
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
return
|
||||
for field in dc_fields:
|
||||
child = getattr(value, field.name)
|
||||
path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
_record_value_paths(child, path, out)
|
||||
return
|
||||
if isinstance(value, Mapping):
|
||||
if not value:
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
return
|
||||
for key, child in value.items():
|
||||
path = f"{prefix}.{key}" if prefix else key
|
||||
_record_value_paths(child, path, out)
|
||||
return
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# __setattr__ patching
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ensure_request_tracking() -> None:
|
||||
for config_type in _TRACKED_REQUEST_TYPES:
|
||||
_patch_tracking_setattr(config_type)
|
||||
|
||||
|
||||
def _patch_tracking_setattr(config_type: type[Any]) -> None:
|
||||
if getattr(config_type, _TRACKING_PATCHED_ATTR, False):
|
||||
return
|
||||
|
||||
original_setattr = cast(
|
||||
Callable[[Any, str, Any], None],
|
||||
config_type.__setattr__,
|
||||
)
|
||||
field_names = {field.name for field in dataclasses.fields(config_type)}
|
||||
|
||||
def _tracking_setattr(self: Any, name: str, value: Any) -> None:
|
||||
if name.startswith("_fastvideo_") or name not in field_names:
|
||||
original_setattr(self, name, value)
|
||||
return
|
||||
|
||||
original_setattr(self, name, value)
|
||||
|
||||
root = getattr(self, _TRACKING_ROOT_ATTR, None)
|
||||
if root is None:
|
||||
return
|
||||
paths = getattr(root, EXPLICIT_PATHS_ATTR, None)
|
||||
if not isinstance(paths, set):
|
||||
return
|
||||
|
||||
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
|
||||
path = f"{prefix}.{name}" if prefix else name
|
||||
# Wholesale dataclass replacement: install roots on the new
|
||||
# instance so its future mutations are tracked too.
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
_set_tracking_roots(root, value, path)
|
||||
_record_value_paths(value, path, paths)
|
||||
|
||||
type.__setattr__(config_type, "__setattr__", _tracking_setattr)
|
||||
setattr(config_type, _TRACKING_PATCHED_ATTR, True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tree walk to set tracking root/path on nested dataclasses
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _set_tracking_roots(
|
||||
root: GenerationRequest,
|
||||
obj: Any,
|
||||
prefix: str,
|
||||
) -> None:
|
||||
if not dataclasses.is_dataclass(obj) or isinstance(obj, type):
|
||||
return
|
||||
object.__setattr__(obj, _TRACKING_ROOT_ATTR, root)
|
||||
object.__setattr__(obj, _TRACKING_PATH_ATTR, prefix)
|
||||
for field in dataclasses.fields(obj):
|
||||
child = getattr(obj, field.name)
|
||||
child_path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
if dataclasses.is_dataclass(child) and not isinstance(child, type):
|
||||
_set_tracking_roots(root, child, child_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EXPLICIT_PATHS_ATTR",
|
||||
"bind_generation_request_raw",
|
||||
"bind_run_config_raw",
|
||||
"bind_serve_config_raw",
|
||||
"get_explicit_paths",
|
||||
"reset_tracking_roots",
|
||||
]
|
||||
@@ -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,10 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass, field, fields
|
||||
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__)
|
||||
|
||||
|
||||
@@ -30,6 +36,16 @@ class SamplingParam:
|
||||
|
||||
# Camera control inputs (HYWorld)
|
||||
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
|
||||
# Camera/action control inputs (GameCraft)
|
||||
camera_states: Any | None = None # Plücker coordinates [B, T_video, 6, H, W]
|
||||
camera_trajectory: str | None = None
|
||||
action_list: list[str] | None = None
|
||||
action_speed_list: list[float] | None = None
|
||||
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
|
||||
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
|
||||
|
||||
# Camera control inputs (LingBotWorld)
|
||||
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
@@ -68,6 +84,7 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
@@ -80,6 +97,54 @@ class SamplingParam:
|
||||
movement_distance: float | None = None
|
||||
camera_rotation: str | None = None
|
||||
|
||||
# LTX-2 multi-modal CFG and STG.
|
||||
# cfg_scale defaults are 1.0 (CFG off) so ``ForwardBatch.__post_init__``
|
||||
# doesn't force ``do_classifier_free_guidance`` on non-LTX-2 models that
|
||||
# never override these fields. LTX-2 presets that need text-CFG on set
|
||||
# them in their ``defaults`` dict (e.g. ``ltx2_base``).
|
||||
ltx2_cfg_scale_video: float = 1.0
|
||||
ltx2_cfg_scale_audio: float = 1.0
|
||||
ltx2_modality_scale_video: float = 3.0
|
||||
ltx2_modality_scale_audio: float = 3.0
|
||||
ltx2_rescale_scale: float = 0.7
|
||||
ltx2_stg_scale_video: float = 1.0
|
||||
ltx2_stg_scale_audio: float = 1.0
|
||||
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
|
||||
@@ -94,26 +159,58 @@ class SamplingParam:
|
||||
raise ValueError("prompt_path must be a txt file")
|
||||
|
||||
def update(self, source_dict: dict[str, Any]) -> None:
|
||||
valid_fields = {f.name for f in fields(self)}
|
||||
for key, value in source_dict.items():
|
||||
if hasattr(self, key):
|
||||
if key in valid_fields:
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
logger.exception("%s has no attribute %s", type(self).__name__, key)
|
||||
logger.error("%s has no field %s", type(self).__name__, key)
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.registry import get_sampling_param_cls_for_name
|
||||
sampling_cls = get_sampling_param_cls_for_name(model_path)
|
||||
if sampling_cls is not None:
|
||||
sampling_param: SamplingParam = sampling_cls()
|
||||
else:
|
||||
logger.warning("Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
model_path)
|
||||
sampling_param = cls()
|
||||
def from_pretrained(cls, model_path: str) -> SamplingParam:
|
||||
sampling_param = cls._from_preset(model_path)
|
||||
if sampling_param is not None:
|
||||
return sampling_param
|
||||
|
||||
return sampling_param
|
||||
logger.warning(
|
||||
"Couldn't find a preset for %s."
|
||||
" Using the default sampling param.",
|
||||
model_path,
|
||||
)
|
||||
return cls()
|
||||
|
||||
@classmethod
|
||||
def _from_preset(
|
||||
cls,
|
||||
model_path: str,
|
||||
) -> SamplingParam | None:
|
||||
"""Build a SamplingParam from preset defaults.
|
||||
|
||||
Returns ``None`` when no preset is configured for
|
||||
*model_path*, letting the caller fall back to the legacy
|
||||
subclass lookup.
|
||||
"""
|
||||
from fastvideo.registry import get_preset_selection
|
||||
|
||||
try:
|
||||
preset_name, model_family = get_preset_selection(model_path)
|
||||
except (ValueError, RuntimeError):
|
||||
return None
|
||||
if preset_name is None or model_family is None:
|
||||
return None
|
||||
|
||||
from fastvideo.api.presets import get_preset
|
||||
|
||||
preset = get_preset(preset_name, model_family)
|
||||
sp = cls()
|
||||
valid_fields = {f.name for f in fields(cls)}
|
||||
for key, value in preset.defaults.items():
|
||||
if key in valid_fields:
|
||||
setattr(sp, key, copy.deepcopy(value))
|
||||
sp.__post_init__()
|
||||
return sp
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: Any) -> Any:
|
||||
+90
-4
@@ -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
|
||||
@@ -73,10 +90,12 @@ class ComponentConfig:
|
||||
@dataclass
|
||||
class PipelineSelection:
|
||||
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
|
||||
profile: str | None = None
|
||||
profile_version: str | None = None
|
||||
preset: str | None = None
|
||||
preset_version: int | None = None
|
||||
components: ComponentConfig = field(default_factory=ComponentConfig)
|
||||
profile_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
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)
|
||||
|
||||
|
||||
@@ -180,11 +199,73 @@ class RunConfig:
|
||||
request: GenerationRequest
|
||||
|
||||
|
||||
@dataclass
|
||||
class WarmupConfig:
|
||||
enabled: bool = True
|
||||
prompt: str = ("A cinematic drone shot over coastal cliffs at sunrise, "
|
||||
"golden light, gentle ocean waves, ultra detailed")
|
||||
timeout_seconds: int = 2400
|
||||
|
||||
|
||||
@dataclass
|
||||
class GpuPoolConfig:
|
||||
num_workers: int | None = None
|
||||
enable_audio_reencode: bool = True
|
||||
conditioning_num_frames: int = 9
|
||||
conditioning_end_offset: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptEnhancerConfig:
|
||||
enabled: bool = False
|
||||
provider: Literal["cerebras", "groq"] = "cerebras"
|
||||
model: str = "gpt-oss-120b"
|
||||
timeout_ms: int = 20000
|
||||
system_prompt_dir: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptSafetyConfig:
|
||||
enabled: bool = False
|
||||
classifier_path: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamingConfig:
|
||||
session_timeout_seconds: int = 300
|
||||
generation_segment_cap: int = 6
|
||||
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
|
||||
warmup: WarmupConfig = field(default_factory=WarmupConfig)
|
||||
pool: GpuPoolConfig = field(default_factory=GpuPoolConfig)
|
||||
prompt: PromptEnhancerConfig = field(default_factory=PromptEnhancerConfig)
|
||||
safety: PromptSafetyConfig = field(default_factory=PromptSafetyConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServeConfig:
|
||||
"""Typed serve config loaded from ``fastvideo serve --config``.
|
||||
|
||||
``default_request`` is a full :class:`GenerationRequest` — the same type
|
||||
clients POST to ``/v1/videos``. At request time the server merges it into
|
||||
the incoming body as the operator-pinned baseline.
|
||||
|
||||
Important nuance: only fields the operator **explicitly wrote** in the
|
||||
serve YAML/JSON count as defaults. Although the in-memory object is
|
||||
fully populated (schema defaults fill every unset field), the merge
|
||||
walks ``_fastvideo_explicit_paths`` — populated during parse — so
|
||||
unset fields are *not* forced onto requests. Per-request precedence:
|
||||
|
||||
body (client-explicit) > default_request (operator-explicit)
|
||||
> hardcoded fallback (e.g. ``fps=24``)
|
||||
|
||||
See :func:`fastvideo.api.compat.explicit_request_updates` for the
|
||||
projection and ``entrypoints/openai/video_api.py::_build_generation_kwargs``
|
||||
for the merge.
|
||||
"""
|
||||
generator: GeneratorConfig
|
||||
server: ServerConfig = field(default_factory=ServerConfig)
|
||||
default_request: GenerationRequest = field(default_factory=GenerationRequest)
|
||||
streaming: StreamingConfig | None = None
|
||||
|
||||
|
||||
__all__ = [
|
||||
@@ -195,16 +276,21 @@ __all__ = [
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"GpuPoolConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"PromptEnhancerConfig",
|
||||
"PromptSafetyConfig",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
"StreamingConfig",
|
||||
"WarmupConfig",
|
||||
]
|
||||
|
||||
@@ -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,13 +0,0 @@
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.hunyuangamecraft import (
|
||||
HunyuanGameCraftSamplingParam,
|
||||
HunyuanGameCraft65FrameSamplingParam,
|
||||
HunyuanGameCraft129FrameSamplingParam,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SamplingParam",
|
||||
"HunyuanGameCraftSamplingParam",
|
||||
"HunyuanGameCraft65FrameSamplingParam",
|
||||
"HunyuanGameCraft129FrameSamplingParam",
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos_Predict2_2B_Video2World_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 93
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = "The video captures a series of frames showing ugly scenes, static with no motion, motion blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of poor quality."
|
||||
num_inference_steps: int = 35
|
||||
@@ -1,23 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25SamplingParamBase(SamplingParam):
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 77
|
||||
fps: int = 24
|
||||
seed: int = 0
|
||||
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, "
|
||||
"low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, "
|
||||
"unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
|
||||
"Overall, the video is of poor quality.")
|
||||
num_inference_steps: int = 35
|
||||
@@ -1,24 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Gen3C_Cosmos_7B_SamplingParam(SamplingParam):
|
||||
"""Defaults for GEN3C (Cosmos-7B) camera-controlled video generation."""
|
||||
|
||||
# Video parameters (matching official GEN3C defaults)
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 121
|
||||
fps: int = 24
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 35
|
||||
|
||||
# GEN3C camera control defaults
|
||||
trajectory_type: str = "left"
|
||||
movement_distance: float = 0.3
|
||||
camera_rotation: str = "center_facing"
|
||||
@@ -1,21 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanSamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastHunyuanSamplingParam(HunyuanSamplingParam):
|
||||
num_inference_steps: int = 6
|
||||
@@ -1,55 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import numpy as np
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_480P_SamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 121
|
||||
height: int = 480
|
||||
width: int = 848
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 6.0
|
||||
sigmas: list[float] | None = field(default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
|
||||
|
||||
negative_prompt: str = ""
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.sigmas = list(np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_480P_StepDistilled_I2V_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
num_inference_steps: int = 12
|
||||
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_720P_Distilled_I2V_SamplingParam(Hunyuan15_720P_SamplingParam):
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_SR_1080P_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
|
||||
num_inference_steps: int = 12
|
||||
num_inference_steps_sr: int = 8
|
||||
|
||||
guidance_scale: float = 1.0
|
||||
@@ -1,92 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Sampling parameters for HunyuanGameCraft video generation.
|
||||
|
||||
GameCraft generates game-like videos with camera/action control.
|
||||
Default parameters are based on the official implementation.
|
||||
"""
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanGameCraftSamplingParam(SamplingParam):
|
||||
"""Sampling parameters for HunyuanGameCraft video generation.
|
||||
|
||||
Supports camera/action conditioning via:
|
||||
- camera_trajectory: Plücker coordinates for camera motion
|
||||
- action_list: List of actions (e.g., ["forward", "left", "right"])
|
||||
- action_speed_list: Speed multipliers for each action
|
||||
|
||||
Default resolution is 704x1280 (same as HunyuanVideo).
|
||||
Default frame count is 33 video frames -> 9 latent frames.
|
||||
"""
|
||||
|
||||
# Number of denoising steps
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Video dimensions
|
||||
# 33 video frames -> 9 latent frames (4x temporal compression)
|
||||
num_frames: int = 33
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
fps: int = 24
|
||||
|
||||
# Guidance scale - official GameCraft uses CFG with guidance_scale=6.0
|
||||
guidance_scale: float = 6.0
|
||||
|
||||
# Negative prompt for CFG (empty string = unconditional)
|
||||
negative_prompt: str = ""
|
||||
|
||||
# Camera/Action conditioning
|
||||
# Camera states as Plücker coordinates [B, T_video, 6, H, W]
|
||||
camera_states: Any | None = None
|
||||
|
||||
# Camera trajectory file/identifier (alternative to camera_states)
|
||||
camera_trajectory: str | None = None
|
||||
|
||||
# Action list for camera motion (e.g., ["forward", "left"])
|
||||
action_list: list[str] | None = None
|
||||
|
||||
# Speed multipliers for each action
|
||||
action_speed_list: list[float] | None = None
|
||||
|
||||
# History frame conditioning (for autoregressive generation)
|
||||
# Ground truth latents for conditioning [B, 16, T, H, W]
|
||||
gt_latents: Any | None = None
|
||||
|
||||
# Mask for conditioning (1=use gt, 0=generate) [B, 1, T, H, W]
|
||||
conditioning_mask: Any | None = None
|
||||
|
||||
# Number of conditioning frames (for autoregressive) - maps to num_cond_frames
|
||||
num_cond_frames: int = 0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# Validate action lists
|
||||
if (self.action_list is not None and self.action_speed_list is not None
|
||||
and len(self.action_list) != len(self.action_speed_list)):
|
||||
raise ValueError(f"action_list length ({len(self.action_list)}) must match "
|
||||
f"action_speed_list length ({len(self.action_speed_list)})")
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanGameCraft65FrameSamplingParam(HunyuanGameCraftSamplingParam):
|
||||
"""Sampling parameters for 65-frame GameCraft generation.
|
||||
|
||||
65 video frames -> 17 latent frames (with first frame as key frame).
|
||||
This is useful for longer video generation.
|
||||
"""
|
||||
num_frames: int = 65
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanGameCraft129FrameSamplingParam(HunyuanGameCraftSamplingParam):
|
||||
"""Sampling parameters for 129-frame GameCraft generation.
|
||||
|
||||
129 video frames -> 33 latent frames.
|
||||
This is the maximum supported by the official implementation.
|
||||
"""
|
||||
num_frames: int = 129
|
||||
@@ -1,25 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
import numpy as np
|
||||
|
||||
|
||||
@dataclass
|
||||
class HYWorld_SamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 125
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
fps: int = 24
|
||||
|
||||
# Camera trajectory: pose string (e.g., 'w-31' means generating [1 + 31] latents) or JSON file path
|
||||
pose: str = 'w-31'
|
||||
|
||||
guidance_scale: float = 6.0
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
sigmas: list[float] | None = field(default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
|
||||
|
||||
negative_prompt: str = ""
|
||||
@@ -1,20 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld_SamplingParam(Wan2_2_I2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 5.0 # high_noise
|
||||
guidance_scale_2: float = 5.0 # low_noise
|
||||
num_inference_steps: int = 70
|
||||
boundary_ratio: float | None = 0.947
|
||||
negative_prompt: str | None = ("画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
|
||||
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
|
||||
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
|
||||
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
|
||||
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
|
||||
"皮肤,肢体,面部特征,汽车,电线")
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
@@ -1,69 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2BaseSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 base one-stage T2V.
|
||||
|
||||
Values follow the official LTX-2 one-stage defaults.
|
||||
Multi-modal CFG params are read by ``LTX2DenoisingStage``.
|
||||
"""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 512
|
||||
width: int = 768
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 40
|
||||
guidance_scale: float = 3.0
|
||||
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
|
||||
negative_prompt: str = ("blurry, out of focus, overexposed, underexposed, low contrast, "
|
||||
"washed out colors, excessive noise, grainy texture, poor lighting, "
|
||||
"flickering, motion blur, distorted proportions, unnatural skin "
|
||||
"tones, deformed facial features, asymmetrical face, missing facial "
|
||||
"features, extra limbs, disfigured hands, wrong hand count, "
|
||||
"artifacts around text, inconsistent perspective, camera shake, "
|
||||
"incorrect depth of field, background too sharp, background clutter, "
|
||||
"distracting reflections, harsh shadows, inconsistent lighting "
|
||||
"direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
|
||||
"wrong gender, exaggerated expressions, wrong gaze direction, "
|
||||
"mismatched lip sync, silent or muted audio, distorted voice, "
|
||||
"robotic voice, echo, background noise, off-sync audio, incorrect "
|
||||
"dialogue, added dialogue, repetitive speech, jittery movement, "
|
||||
"awkward pauses, incorrect timing, unnatural transitions, "
|
||||
"inconsistent framing, tilted camera, flat lighting, inconsistent "
|
||||
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
|
||||
# Official LTX-2 multi-modal CFG defaults.
|
||||
ltx2_cfg_scale_video: float = 3.0
|
||||
ltx2_cfg_scale_audio: float = 7.0
|
||||
ltx2_modality_scale_video: float = 3.0
|
||||
ltx2_modality_scale_audio: float = 3.0
|
||||
ltx2_rescale_scale: float = 0.7
|
||||
# STG (Spatio-Temporal Guidance) defaults from official LTX-2.
|
||||
ltx2_stg_scale_video: float = 1.0
|
||||
ltx2_stg_scale_audio: float = 1.0
|
||||
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
|
||||
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2DistilledSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 1024
|
||||
width: int = 1536
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 8
|
||||
guidance_scale: float = 1.0
|
||||
# No default negative_prompt for distilled models
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
# Backward compatibility alias.
|
||||
LTX2SamplingParam = LTX2DistilledSamplingParam
|
||||
@@ -1,25 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class SD35SamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 512
|
||||
width: int = 512
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 6.0
|
||||
@@ -1,73 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
TurboDiffusion sampling parameters.
|
||||
|
||||
TurboDiffusion uses RCM (recurrent Consistency Model) scheduler for
|
||||
1-4 step video generation with no classifier-free guidance.
|
||||
"""
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class TurboDiffusionT2V_1_3B_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for TurboDiffusion T2V 1.3B model.
|
||||
|
||||
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
|
||||
"""
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
# No negative prompt needed for TurboDiffusion (no CFG)
|
||||
negative_prompt: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TurboDiffusionT2V_14B_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for TurboDiffusion T2V 14B model.
|
||||
|
||||
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
|
||||
"""
|
||||
# Video parameters (720p for 14B)
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
# No negative prompt needed for TurboDiffusion (no CFG)
|
||||
negative_prompt: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TurboDiffusionI2V_A14B_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for TurboDiffusion I2V A14B model.
|
||||
|
||||
Uses 4-step RCM sampling with dual-model switching (high/low noise).
|
||||
"""
|
||||
# Video parameters (720p for A14B I2V)
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 4
|
||||
|
||||
# Note: boundary_ratio is set in the pipeline config (TurboDiffusionI2VConfig),
|
||||
# not here. This keeps sampling params and pipeline config separate.
|
||||
|
||||
# No negative prompt needed for TurboDiffusion (no CFG)
|
||||
negative_prompt: str | None = None
|
||||
@@ -1,154 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V_1_3B_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V_14B_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V_14B_480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# DMD parameters
|
||||
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
|
||||
num_inference_steps: int = 3
|
||||
num_frames: int = 61
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_Control_SamplingParam(SamplingParam):
|
||||
fps: int = 16
|
||||
num_frames: int = 49
|
||||
height: int = 832
|
||||
width: int = 480
|
||||
guidance_scale: float = 6.0
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.2 TI2V Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_2_Base_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 121
|
||||
fps: int = 24
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 4.0 # high_noise
|
||||
guidance_scale_2: float = 3.0 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 3.5 # high_noise
|
||||
guidance_scale_2: float = 3.5 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_Fun_A14B_Control_SamplingParam(Wan2_1_Fun_1_3B_Control_SamplingParam):
|
||||
num_frames: int = 81
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Causal Self-Forcing =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(Wan2_1_Fun_1_3B_InP_SamplingParam):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(Wan2_2_T2V_A14B_SamplingParam):
|
||||
num_inference_steps: int = 8
|
||||
num_frames: int = 81
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatrixGame2_SamplingParam(SamplingParam):
|
||||
height: int = 352
|
||||
width: int = 640
|
||||
num_frames: int = 57
|
||||
fps: int = 25
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 3
|
||||
negative_prompt: str | None = None
|
||||
@@ -7,7 +7,7 @@ Example usage:
|
||||
# launch a server and benchmark on it
|
||||
|
||||
# T2V or T2I or any other multimodal generation model
|
||||
fastvideo serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers --port 8000
|
||||
fastvideo serve --config serve.yaml
|
||||
|
||||
# benchmark it and make sure the port is the same as the server's port
|
||||
fastvideo bench --dataset vbench --num-prompts 20 --port 8000
|
||||
|
||||
@@ -2,19 +2,17 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.utils import RaiseNotImplementedAction
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.entrypoints.cli.inference_config import build_generate_run_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_VALIDATED_RUN_CONFIG_ATTR = "_fastvideo_validated_run_config"
|
||||
|
||||
|
||||
class GenerateSubcommand(CLISubcommand):
|
||||
@@ -23,89 +21,47 @@ class GenerateSubcommand(CLISubcommand):
|
||||
def __init__(self) -> None:
|
||||
self.name = "generate"
|
||||
super().__init__()
|
||||
self.init_arg_names = self._get_init_arg_names()
|
||||
self.generation_arg_names = self._get_generation_arg_names()
|
||||
|
||||
def _get_init_arg_names(self) -> list[str]:
|
||||
"""Get names of arguments for VideoGenerator initialization"""
|
||||
return ["num_gpus", "tp_size", "sp_size", "model_path"]
|
||||
|
||||
def _get_generation_arg_names(self) -> list[str]:
|
||||
"""Get names of arguments for generate_video method"""
|
||||
return [field.name for field in dataclasses.fields(SamplingParam)]
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = ['subparser', 'config', 'dispatch_function']
|
||||
run_config = getattr(args, _VALIDATED_RUN_CONFIG_ATTR, None)
|
||||
if run_config is None:
|
||||
run_config = build_generate_run_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
)
|
||||
logger.info("CLI generate config: %s", run_config)
|
||||
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (k not in excluded_args and v is not None and hasattr(args, '_provided') and k in args._provided):
|
||||
provided_args[k] = v
|
||||
|
||||
if 'model_path' in vars(args) and args.model_path is not None:
|
||||
provided_args['model_path'] = args.model_path
|
||||
|
||||
if 'prompt' in vars(args) and args.prompt is not None:
|
||||
provided_args['prompt'] = args.prompt
|
||||
|
||||
merged_args = {**provided_args}
|
||||
|
||||
logger.info('CLI Args: %s', merged_args)
|
||||
|
||||
if 'model_path' not in merged_args or not merged_args['model_path']:
|
||||
raise ValueError("model_path must be provided either in config file or via --model-path")
|
||||
|
||||
# Check if either prompt or prompt_txt is provided
|
||||
has_prompt = 'prompt' in merged_args and merged_args['prompt']
|
||||
has_prompt_txt = 'prompt_txt' in merged_args and merged_args['prompt_txt']
|
||||
|
||||
if not (has_prompt or has_prompt_txt):
|
||||
raise ValueError("Either prompt or prompt_txt must be provided")
|
||||
|
||||
if has_prompt and has_prompt_txt:
|
||||
raise ValueError("Cannot provide both 'prompt' and 'prompt_txt'. Use only one of them.")
|
||||
|
||||
init_args = {k: v for k, v in merged_args.items() if k not in self.generation_arg_names}
|
||||
generation_args = {k: v for k, v in merged_args.items() if k in self.generation_arg_names}
|
||||
generation_args.setdefault("return_frames", False)
|
||||
|
||||
model_path = init_args.pop('model_path')
|
||||
prompt = generation_args.pop('prompt', None)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_path=model_path, **init_args)
|
||||
|
||||
# Call generate_video - it handles both single and batch modes
|
||||
generator.generate_video(prompt=prompt, **generation_args)
|
||||
generator = VideoGenerator.from_config(run_config.generator)
|
||||
generator.generate(run_config.request)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
"""Validate the arguments for this command"""
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
if args.config and not os.path.exists(args.config):
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo generate requires --config PATH; use a nested "
|
||||
"run config plus optional dotted overrides")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Config file not found: {args.config}")
|
||||
setattr(
|
||||
args,
|
||||
_VALIDATED_RUN_CONFIG_ATTR,
|
||||
build_generate_run_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
),
|
||||
)
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
generate_parser = subparsers.add_parser(
|
||||
"generate",
|
||||
help="Run inference on a model",
|
||||
usage="fastvideo generate (--model-path MODEL_PATH_OR_ID --prompt PROMPT) | --config CONFIG_FILE [OPTIONS]")
|
||||
usage="fastvideo generate --config RUN_CONFIG [--dotted.override VALUE]")
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default='',
|
||||
required=False,
|
||||
help="Read CLI options from a config JSON or YAML file. If provided, --model-path and --prompt are optional."
|
||||
)
|
||||
|
||||
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
|
||||
generate_parser = SamplingParam.add_cli_args(generate_parser)
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--text-encoder-configs",
|
||||
action=RaiseNotImplementedAction,
|
||||
help="JSON array of text encoder configurations (NOT YET IMPLEMENTED)",
|
||||
help="Path to a nested run config JSON or YAML file. Required.",
|
||||
)
|
||||
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
|
||||
from fastvideo.api.parser import load_raw_config, parse_config
|
||||
from fastvideo.api.schema import RunConfig, ServeConfig
|
||||
|
||||
_GENERATE_OVERRIDE_PREFIXES = ("generator.", "request.")
|
||||
_SERVE_OVERRIDE_PREFIXES = (
|
||||
"generator.",
|
||||
"server.",
|
||||
"default_request.",
|
||||
)
|
||||
|
||||
|
||||
def build_generate_run_config(
|
||||
args: argparse.Namespace,
|
||||
overrides: list[str] | None = None,
|
||||
) -> RunConfig:
|
||||
raw = _load_nested_config(getattr(args, "config", None))
|
||||
raw.setdefault("request", {})
|
||||
raw = _apply_dotted_overrides(
|
||||
raw,
|
||||
overrides,
|
||||
allowed_prefixes=_GENERATE_OVERRIDE_PREFIXES,
|
||||
)
|
||||
_ensure_generate_cli_defaults(raw)
|
||||
config = parse_config(RunConfig, raw)
|
||||
_validate_num_gpus(config.generator.engine.num_gpus)
|
||||
_validate_generate_prompt_sources(config)
|
||||
return config
|
||||
|
||||
|
||||
def build_serve_config(
|
||||
args: argparse.Namespace,
|
||||
overrides: list[str] | None = None,
|
||||
) -> ServeConfig:
|
||||
raw = _load_nested_config(getattr(args, "config", None))
|
||||
raw.setdefault("server", {})
|
||||
raw.setdefault("default_request", {})
|
||||
raw = _apply_dotted_overrides(
|
||||
raw,
|
||||
overrides,
|
||||
allowed_prefixes=_SERVE_OVERRIDE_PREFIXES,
|
||||
)
|
||||
config = parse_config(ServeConfig, raw)
|
||||
_validate_num_gpus(config.generator.engine.num_gpus)
|
||||
return config
|
||||
|
||||
|
||||
def _load_nested_config(path: str | None) -> dict[str, Any]:
|
||||
if not path:
|
||||
raise ValueError("Inference CLI requires --config PATH; use a nested config file "
|
||||
"plus optional dotted overrides")
|
||||
|
||||
raw = load_raw_config(path)
|
||||
if not isinstance(raw.get("generator"), Mapping):
|
||||
raise ValueError("Inference config must use the nested schema with a top-level "
|
||||
"'generator' mapping")
|
||||
return deepcopy(dict(raw))
|
||||
|
||||
|
||||
def _apply_dotted_overrides(
|
||||
raw: Mapping[str, Any],
|
||||
overrides: list[str] | None,
|
||||
*,
|
||||
allowed_prefixes: tuple[str, ...],
|
||||
) -> dict[str, Any]:
|
||||
if not overrides:
|
||||
return deepcopy(dict(raw))
|
||||
|
||||
parsed = parse_cli_overrides(overrides)
|
||||
for key in parsed:
|
||||
if "." not in key:
|
||||
raise ValueError("CLI overrides must use dotted config paths like "
|
||||
"--request.sampling.seed 42")
|
||||
if not key.startswith(allowed_prefixes):
|
||||
allowed = ", ".join(allowed_prefixes)
|
||||
raise ValueError(f"Unsupported override path {key!r}. Allowed prefixes: {allowed}")
|
||||
return apply_overrides(raw, parsed)
|
||||
|
||||
|
||||
def _ensure_generate_cli_defaults(raw: dict[str, Any]) -> None:
|
||||
request = raw.setdefault("request", {})
|
||||
output = request.setdefault("output", {})
|
||||
output.setdefault("return_frames", False)
|
||||
|
||||
|
||||
def _validate_generate_prompt_sources(config: RunConfig) -> None:
|
||||
has_prompt = config.request.prompt is not None
|
||||
has_prompt_path = config.request.inputs.prompt_path is not None
|
||||
if not (has_prompt or has_prompt_path):
|
||||
raise ValueError("Either request.prompt or request.inputs.prompt_path must be provided")
|
||||
if has_prompt and has_prompt_path:
|
||||
raise ValueError("Cannot provide both request.prompt and request.inputs.prompt_path")
|
||||
|
||||
|
||||
def _validate_num_gpus(num_gpus: int) -> None:
|
||||
if num_gpus <= 0:
|
||||
raise ValueError(f"generator.engine.num_gpus must be > 0; got {num_gpus}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_generate_run_config",
|
||||
"build_serve_config",
|
||||
]
|
||||
@@ -1,6 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
@@ -27,14 +26,17 @@ def main() -> None:
|
||||
for cmd in cmd_init():
|
||||
cmd.subparser_init(subparsers).set_defaults(dispatch_function=cmd.cmd)
|
||||
cmds[cmd.name] = cmd
|
||||
args = parser.parse_args()
|
||||
|
||||
args, unknown = parser.parse_known_args()
|
||||
if unknown and args.subparser not in {"generate", "serve"}:
|
||||
parser.error(f"unrecognized arguments: {' '.join(unknown)}")
|
||||
args._unknown = unknown
|
||||
if args.subparser in cmds:
|
||||
cmds[args.subparser].validate(args)
|
||||
|
||||
if hasattr(args, "dispatch_function"):
|
||||
args.dispatch_function(args)
|
||||
else:
|
||||
parser.print_help()
|
||||
return
|
||||
|
||||
parser.print_help()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -2,14 +2,17 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.api.compat import generator_config_to_fastvideo_args
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.entrypoints.cli.inference_config import build_serve_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_VALIDATED_SERVE_CONFIG_ATTR = "_fastvideo_validated_serve_config"
|
||||
|
||||
|
||||
class ServeSubcommand(CLISubcommand):
|
||||
@@ -20,94 +23,69 @@ class ServeSubcommand(CLISubcommand):
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = {
|
||||
"subparser",
|
||||
"config",
|
||||
"dispatch_function",
|
||||
"host",
|
||||
"port",
|
||||
"output_dir",
|
||||
}
|
||||
serve_config = getattr(args, _VALIDATED_SERVE_CONFIG_ATTR, None)
|
||||
if serve_config is None:
|
||||
serve_config = build_serve_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
)
|
||||
|
||||
provided: set[str] = getattr(args, '_provided', set())
|
||||
cli_kwargs = {}
|
||||
for k, v in vars(args).items():
|
||||
if k in excluded_args:
|
||||
continue
|
||||
if k == '_provided':
|
||||
continue
|
||||
if k in provided and v is not None:
|
||||
cli_kwargs[k] = v
|
||||
logger.info("CLI serve config: %s", serve_config)
|
||||
|
||||
if 'model_path' not in cli_kwargs and args.model_path is not None:
|
||||
cli_kwargs['model_path'] = args.model_path
|
||||
|
||||
if not cli_kwargs.get('model_path'):
|
||||
raise ValueError("model_path must be provided via --model-path")
|
||||
# A `streaming:` block selects the WebSocket/Dynamo runtime;
|
||||
# its deps stay out of REST-only deployments via lazy import.
|
||||
if serve_config.streaming is not None:
|
||||
from fastvideo.entrypoints.streaming.server import (
|
||||
run_server as run_streaming_server, )
|
||||
run_streaming_server(serve_config)
|
||||
return
|
||||
|
||||
from fastvideo.entrypoints.openai.api_server import (
|
||||
DEFAULT_HOST,
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
DEFAULT_PORT,
|
||||
run_server,
|
||||
run_server, )
|
||||
|
||||
logger.info(
|
||||
"Server will listen on %s:%d",
|
||||
serve_config.server.host,
|
||||
serve_config.server.port,
|
||||
)
|
||||
|
||||
host = getattr(args, "host", DEFAULT_HOST)
|
||||
port = getattr(args, "port", DEFAULT_PORT)
|
||||
output_dir = getattr(args, "output_dir", DEFAULT_OUTPUT_DIR)
|
||||
|
||||
logger.info("CLI serve args: %s", cli_kwargs)
|
||||
logger.info("Server will listen on %s:%d", host, port)
|
||||
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(**cli_kwargs)
|
||||
run_server(fastvideo_args, host=host, port=port, output_dir=output_dir)
|
||||
fastvideo_args = generator_config_to_fastvideo_args(serve_config.generator)
|
||||
run_server(
|
||||
fastvideo_args,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
output_dir=serve_config.server.output_dir,
|
||||
default_request=serve_config.default_request,
|
||||
)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
from fastvideo.entrypoints.openai.api_server import (
|
||||
DEFAULT_HOST,
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
DEFAULT_PORT,
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo serve requires --config PATH; use a nested "
|
||||
"serve config plus optional dotted overrides")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Config file not found: {args.config}")
|
||||
setattr(
|
||||
args,
|
||||
_VALIDATED_SERVE_CONFIG_ATTR,
|
||||
build_serve_config(
|
||||
args,
|
||||
overrides=getattr(args, "_unknown", None),
|
||||
),
|
||||
)
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
serve_parser = subparsers.add_parser(
|
||||
"serve",
|
||||
help="Start an OpenAI-compatible HTTP server",
|
||||
usage=("fastvideo serve --model-path MODEL_PATH_OR_ID "
|
||||
"[--host HOST] [--port PORT] [OPTIONS]"),
|
||||
)
|
||||
|
||||
serve_parser.add_argument(
|
||||
"--host",
|
||||
type=str,
|
||||
default=DEFAULT_HOST,
|
||||
help=f"Host to bind the server to (default: {DEFAULT_HOST})",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--port",
|
||||
type=int,
|
||||
default=DEFAULT_PORT,
|
||||
help=f"Port to listen on (default: {DEFAULT_PORT})",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--output-dir",
|
||||
type=str,
|
||||
default=DEFAULT_OUTPUT_DIR,
|
||||
help=("Directory for generated outputs "
|
||||
f"(default: {DEFAULT_OUTPUT_DIR})"),
|
||||
usage="fastvideo serve --config SERVE_CONFIG [--dotted.override VALUE]",
|
||||
)
|
||||
serve_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="Read CLI options from a config JSON or YAML file.",
|
||||
help="Path to a nested config JSON or YAML file. Required.",
|
||||
)
|
||||
|
||||
serve_parser = FastVideoArgs.add_cli_args(serve_parser)
|
||||
return cast(FlexibleArgumentParser, serve_parser)
|
||||
|
||||
|
||||
|
||||
@@ -8,6 +8,8 @@ import uvicorn
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from fastvideo.api.presets import validate_preset_selection
|
||||
from fastvideo.api.schema import GenerationRequest
|
||||
from fastvideo.entrypoints.openai.state import (
|
||||
DEFAULT_OUTPUT_DIR,
|
||||
clear_state,
|
||||
@@ -16,6 +18,7 @@ from fastvideo.entrypoints.openai.state import (
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.registry import get_preset_selection
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -23,17 +26,40 @@ DEFAULT_HOST = "0.0.0.0"
|
||||
DEFAULT_PORT = 8000
|
||||
|
||||
|
||||
def _validate_default_request_against_preset(
|
||||
default_request: GenerationRequest,
|
||||
model_path: str,
|
||||
) -> None:
|
||||
"""Validate ``default_request.stage_overrides`` against the model's preset.
|
||||
|
||||
Called once at server startup from :func:`run_server`. The
|
||||
``default_request`` is static server config, so validation results are
|
||||
invariant across requests — there's no reason to re-run per request.
|
||||
"""
|
||||
if not default_request.stage_overrides:
|
||||
return
|
||||
preset_name, model_family = get_preset_selection(model_path)
|
||||
if preset_name is None or model_family is None:
|
||||
return
|
||||
validate_preset_selection(
|
||||
preset_name,
|
||||
model_family,
|
||||
stage_overrides=default_request.stage_overrides,
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
||||
"""Load model on startup, clean up on shutdown"""
|
||||
args: FastVideoArgs = app.state.fastvideo_args
|
||||
output_dir: str = app.state.output_dir
|
||||
default_request: GenerationRequest | None = getattr(app.state, "default_request", None)
|
||||
|
||||
logger.info("Loading model from %s ...", args.model_path)
|
||||
generator = VideoGenerator.from_fastvideo_args(args)
|
||||
logger.info("Model loaded successfully.")
|
||||
|
||||
set_state(generator, args, output_dir)
|
||||
set_state(generator, args, output_dir, default_request=default_request)
|
||||
|
||||
yield # server is running
|
||||
|
||||
@@ -46,6 +72,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
||||
def create_app(
|
||||
fastvideo_args: FastVideoArgs,
|
||||
output_dir: str = DEFAULT_OUTPUT_DIR,
|
||||
default_request: GenerationRequest | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI application with all routers mounted"""
|
||||
|
||||
@@ -56,6 +83,7 @@ def create_app(
|
||||
)
|
||||
app.state.fastvideo_args = fastvideo_args
|
||||
app.state.output_dir = output_dir
|
||||
app.state.default_request = default_request
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
@@ -108,9 +136,17 @@ def run_server(
|
||||
host: str = DEFAULT_HOST,
|
||||
port: int = DEFAULT_PORT,
|
||||
output_dir: str = DEFAULT_OUTPUT_DIR,
|
||||
default_request: GenerationRequest | None = None,
|
||||
):
|
||||
"""Create the app and run it with uvicorn"""
|
||||
app = create_app(fastvideo_args, output_dir=output_dir)
|
||||
if default_request is not None:
|
||||
_validate_default_request_against_preset(default_request, fastvideo_args.model_path)
|
||||
|
||||
app = create_app(
|
||||
fastvideo_args,
|
||||
output_dir=output_dir,
|
||||
default_request=default_request,
|
||||
)
|
||||
|
||||
logger.info("Starting FastVideo server on %s:%d", host, port)
|
||||
logger.info("Model: %s", fastvideo_args.model_path)
|
||||
|
||||
@@ -10,6 +10,7 @@ from __future__ import annotations
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.api.schema import GenerationRequest
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
|
||||
@@ -18,6 +19,7 @@ DEFAULT_OUTPUT_DIR = "outputs"
|
||||
_generator: VideoGenerator | None = None
|
||||
_fastvideo_args: FastVideoArgs | None = None
|
||||
_output_dir: str = DEFAULT_OUTPUT_DIR
|
||||
_default_request: GenerationRequest | None = None
|
||||
|
||||
|
||||
def get_generator() -> VideoGenerator:
|
||||
@@ -37,20 +39,28 @@ def get_output_dir() -> str:
|
||||
return _output_dir
|
||||
|
||||
|
||||
def get_default_request() -> GenerationRequest | None:
|
||||
"""Return the ServeConfig.default_request set at startup, if any."""
|
||||
return _default_request
|
||||
|
||||
|
||||
def set_state(
|
||||
generator: VideoGenerator,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
output_dir: str,
|
||||
default_request: GenerationRequest | None = None,
|
||||
) -> None:
|
||||
"""Set all server state at once (called from lifespan)."""
|
||||
global _generator, _fastvideo_args, _output_dir
|
||||
global _generator, _fastvideo_args, _output_dir, _default_request
|
||||
_generator = generator
|
||||
_fastvideo_args = fastvideo_args
|
||||
_output_dir = output_dir
|
||||
_default_request = default_request
|
||||
|
||||
|
||||
def clear_state() -> None:
|
||||
"""Clear server state on shutdown."""
|
||||
global _generator, _fastvideo_args
|
||||
global _generator, _fastvideo_args, _default_request
|
||||
_generator = None
|
||||
_fastvideo_args = None
|
||||
_default_request = None
|
||||
|
||||
@@ -19,7 +19,10 @@ from fastapi import (
|
||||
)
|
||||
from fastapi.responses import FileResponse
|
||||
|
||||
from fastvideo.api.compat import explicit_request_updates
|
||||
from fastvideo.api.schema import GenerationRequest
|
||||
from fastvideo.entrypoints.openai.state import (
|
||||
get_default_request,
|
||||
get_generator,
|
||||
get_output_dir,
|
||||
get_server_args,
|
||||
@@ -42,49 +45,73 @@ logger = init_logger(__name__)
|
||||
router = APIRouter(prefix="/v1/videos", tags=["videos"])
|
||||
|
||||
|
||||
def _build_generation_kwargs(request_id: str, req: VideoGenerationsRequest) -> dict[str, Any]:
|
||||
def _build_generation_kwargs(
|
||||
request_id: str,
|
||||
req: VideoGenerationsRequest,
|
||||
default_request: GenerationRequest | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build a flat kwargs dict for ``generator.generate_video``.
|
||||
|
||||
Precedence (highest to lowest):
|
||||
1. Request body — only fields the client explicitly sent
|
||||
(``req.model_fields_set``, Pydantic v2).
|
||||
2. ``default_request`` — only fields the operator explicitly set in
|
||||
the serve YAML, projected via ``explicit_request_updates``. Schema
|
||||
defaults on the dataclass are *not* treated as defaults here.
|
||||
3. Hardcoded fallback (e.g. ``fps=24`` when neither side set it).
|
||||
|
||||
Why gate on ``model_fields_set`` / explicit paths? Both the request
|
||||
Pydantic model and the ``GenerationRequest`` dataclass carry schema
|
||||
defaults (e.g. ``seed=1024``, ``num_frames=125``). Without the gate
|
||||
those would masquerade as intent and shadow the other side — the
|
||||
gate preserves "operator pinned it" vs. "dataclass happened to have
|
||||
that default."
|
||||
"""
|
||||
kwargs: dict[str, Any] = {}
|
||||
if default_request is not None:
|
||||
kwargs.update(explicit_request_updates(default_request))
|
||||
|
||||
body_set = req.model_fields_set
|
||||
kwargs["prompt"] = req.prompt
|
||||
|
||||
# Resolution
|
||||
if req.size:
|
||||
if "size" in body_set and req.size:
|
||||
w, h = parse_size(req.size)
|
||||
if w is not None and h is not None:
|
||||
kwargs["width"] = w
|
||||
kwargs["height"] = h
|
||||
|
||||
# Frame count / duration
|
||||
fps = req.fps if req.fps is not None else 24
|
||||
kwargs["fps"] = fps
|
||||
if "fps" in body_set and req.fps is not None:
|
||||
kwargs["fps"] = req.fps
|
||||
|
||||
if req.num_frames is not None:
|
||||
if "num_frames" in body_set and req.num_frames is not None:
|
||||
kwargs["num_frames"] = req.num_frames
|
||||
elif req.seconds is not None:
|
||||
elif "seconds" in body_set and req.seconds is not None:
|
||||
fps = kwargs.get("fps", 24)
|
||||
kwargs["num_frames"] = fps * req.seconds
|
||||
|
||||
# Sampling parameters
|
||||
if req.seed is not None:
|
||||
if "seed" in body_set and req.seed is not None:
|
||||
kwargs["seed"] = req.seed
|
||||
if req.num_inference_steps is not None:
|
||||
if ("num_inference_steps" in body_set and req.num_inference_steps is not None):
|
||||
kwargs["num_inference_steps"] = req.num_inference_steps
|
||||
if req.guidance_scale is not None:
|
||||
if "guidance_scale" in body_set and req.guidance_scale is not None:
|
||||
kwargs["guidance_scale"] = req.guidance_scale
|
||||
if req.guidance_scale_2 is not None:
|
||||
if "guidance_scale_2" in body_set and req.guidance_scale_2 is not None:
|
||||
kwargs["guidance_scale_2"] = req.guidance_scale_2
|
||||
if req.negative_prompt is not None:
|
||||
if "negative_prompt" in body_set and req.negative_prompt is not None:
|
||||
kwargs["negative_prompt"] = req.negative_prompt
|
||||
if req.enable_teacache:
|
||||
if "enable_teacache" in body_set and req.enable_teacache:
|
||||
kwargs["enable_teacache"] = True
|
||||
if req.true_cfg_scale is not None:
|
||||
if "true_cfg_scale" in body_set and req.true_cfg_scale is not None:
|
||||
kwargs["true_cfg_scale"] = req.true_cfg_scale
|
||||
|
||||
# Image-to-video input
|
||||
if req.input_reference is not None:
|
||||
if "input_reference" in body_set and req.input_reference is not None:
|
||||
kwargs["image_path"] = req.input_reference
|
||||
|
||||
# Output path
|
||||
output_dir = req.output_path or os.path.join(get_output_dir(), "videos")
|
||||
kwargs.setdefault("fps", 24)
|
||||
|
||||
default_output_path = kwargs.pop("output_path", None)
|
||||
body_output_dir = req.output_path if "output_path" in body_set else None
|
||||
output_dir = body_output_dir or default_output_path or os.path.join(get_output_dir(), "videos")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
kwargs["output_path"] = os.path.join(output_dir, f"{request_id}.mp4")
|
||||
kwargs["save_video"] = True
|
||||
@@ -272,7 +299,12 @@ async def create_video(
|
||||
|
||||
logger.info("Video generation request %s: prompt=%s", request_id, req.prompt[:100])
|
||||
|
||||
gen_kwargs = _build_generation_kwargs(request_id, req)
|
||||
# default_request was validated at server startup (run_server) and is
|
||||
# read-only on the request hot path — _build_generation_kwargs and
|
||||
# explicit_request_updates only read, so no per-request deepcopy needed.
|
||||
default_request = get_default_request()
|
||||
|
||||
gen_kwargs = _build_generation_kwargs(request_id, req, default_request=default_request)
|
||||
job = _make_video_job(request_id, req, gen_kwargs)
|
||||
await VIDEO_STORE.upsert(request_id, job)
|
||||
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.entrypoints.streaming.server import build_app, run_server
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
BlobStore,
|
||||
InMemoryBlobStore,
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
|
||||
__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",
|
||||
]
|
||||
@@ -0,0 +1,531 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-generator FastAPI + WebSocket streaming server."""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
AutoExtensionUpdated,
|
||||
ContinuationStateSnapshot,
|
||||
EnhancementUpdated,
|
||||
ErrorMessage,
|
||||
GenerationPausedUpdated,
|
||||
GpuAssigned,
|
||||
LoopGenerationUpdated,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamComplete,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SeedPromptsUpdated,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image, )
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# RFC 6455 WebSocket close codes used by the server.
|
||||
_WS_CLOSE_UNSUPPORTED_DATA = 1003
|
||||
_WS_CLOSE_TRY_AGAIN_LATER = 1013
|
||||
|
||||
|
||||
class _GeneratorProto(Protocol):
|
||||
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
|
||||
|
||||
def generate(self, request: GenerationRequest) -> Any:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerState:
|
||||
serve_config: ServeConfig
|
||||
generator: _GeneratorProto
|
||||
sessions: SessionManager
|
||||
session_store: SessionStore
|
||||
|
||||
|
||||
def build_app(
|
||||
serve_config: ServeConfig,
|
||||
generator: _GeneratorProto,
|
||||
*,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI app used by :func:`run_server`.
|
||||
|
||||
Exposed so tests can drive the WebSocket endpoint in-process via
|
||||
``starlette.testclient.TestClient(app).websocket_connect(...)``.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
|
||||
"server; got None. Add a `streaming:` block to your serve config.")
|
||||
|
||||
sessions = SessionManager(
|
||||
segment_cap=serve_config.streaming.generation_segment_cap,
|
||||
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
|
||||
)
|
||||
state = ServerState(
|
||||
serve_config=serve_config,
|
||||
generator=generator,
|
||||
sessions=sessions,
|
||||
session_store=session_store or InMemorySessionStore(),
|
||||
)
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming")
|
||||
|
||||
@app.get("/health")
|
||||
async def _health() -> JSONResponse:
|
||||
return JSONResponse({
|
||||
"status": "ok",
|
||||
"sessions": len(state.sessions),
|
||||
"stream_mode": state.serve_config.streaming.stream_mode,
|
||||
})
|
||||
|
||||
@app.websocket("/v1/stream")
|
||||
async def _stream(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
try:
|
||||
session = state.sessions.create()
|
||||
except SessionRejected as exc:
|
||||
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
|
||||
return
|
||||
|
||||
try:
|
||||
await _handle_session(websocket, session, state)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("session %s: client disconnected", session.id[:8])
|
||||
except Exception: # pragma: no cover - defensive catch-all
|
||||
logger.exception("session %s: unhandled error", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
finally:
|
||||
_cleanup_session(session, state)
|
||||
|
||||
app.state.server_state = state
|
||||
return app
|
||||
|
||||
|
||||
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
|
||||
"""Launch the streaming server.
|
||||
|
||||
Boots a :class:`fastvideo.VideoGenerator` from
|
||||
``serve_config.generator`` unless ``generator`` is provided, then
|
||||
serves ``build_app(...)`` via uvicorn.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
|
||||
"got None. Add a `streaming:` block to your serve config.")
|
||||
|
||||
import uvicorn
|
||||
|
||||
if generator is None:
|
||||
from fastvideo import VideoGenerator # lazy to avoid boot cost
|
||||
|
||||
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
|
||||
app = build_app(serve_config, generator)
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
)
|
||||
|
||||
|
||||
async def _handle_session(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
init = await _read_init_message(websocket, session, state)
|
||||
if init is None:
|
||||
return
|
||||
|
||||
await _apply_session_init(session, init, state)
|
||||
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
|
||||
session.transition(SessionState.GPU_BINDING)
|
||||
await _send_json(websocket, GpuAssigned(
|
||||
gpu_id=0,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
session.transition(SessionState.ACTIVE)
|
||||
await _send_json(websocket, _build_stream_start(session, state))
|
||||
|
||||
try:
|
||||
await _run_segment_loop(websocket, session, state)
|
||||
finally:
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
|
||||
|
||||
|
||||
async def _read_init_message(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> SessionInitV2 | None:
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: init timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return None
|
||||
except WebSocketDisconnect:
|
||||
return None
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
|
||||
return None
|
||||
if not isinstance(parsed, SessionInitV2):
|
||||
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
|
||||
return None
|
||||
return parsed
|
||||
|
||||
|
||||
async def _reject_init(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
message: str,
|
||||
close_reason: str,
|
||||
) -> None:
|
||||
await _send_error(websocket, "invalid_message", message, retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.REJECTED)
|
||||
|
||||
|
||||
async def _apply_session_init(
|
||||
session: Session,
|
||||
init: SessionInitV2,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
session.client_id = init.client_id
|
||||
session.preset = init.preset
|
||||
session.preset_label = init.preset_label
|
||||
session.curated_prompts = list(init.curated_prompts)
|
||||
session.enhancement_enabled = init.enhancement_enabled
|
||||
session.auto_extension_enabled = init.auto_extension_enabled
|
||||
session.loop_generation_enabled = init.loop_generation_enabled
|
||||
session.single_clip_mode = init.single_clip_mode
|
||||
session.stream_mode = init.stream_mode
|
||||
|
||||
if init.initial_image is not None:
|
||||
# Decode + disk write off the event loop; payload is up to 32 MiB.
|
||||
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
|
||||
if image is not None:
|
||||
session.metadata["session_init_image"] = image.path
|
||||
|
||||
if init.continuation_state is not None:
|
||||
session.continuation_state = _coerce_state(init.continuation_state)
|
||||
if session.continuation_state is not None:
|
||||
state.session_store.store(session.id, session.continuation_state)
|
||||
|
||||
|
||||
async def _run_segment_loop(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
cap = state.sessions.segment_cap
|
||||
while True:
|
||||
if session.segment_cap_reached(cap):
|
||||
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
|
||||
return
|
||||
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: idle timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return
|
||||
except WebSocketDisconnect:
|
||||
return
|
||||
session.touch()
|
||||
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SnapshotState):
|
||||
snap = state.session_store.snapshot(session.id)
|
||||
if snap is None:
|
||||
await _send_error(websocket,
|
||||
"internal_error",
|
||||
"no continuation state available for session",
|
||||
retryable=False)
|
||||
continue
|
||||
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SegmentPromptSource):
|
||||
await _run_segment(websocket, session, state, parsed)
|
||||
continue
|
||||
|
||||
# Silently ignore unknown-but-valid types (additive-evolution
|
||||
# rule in streaming.md).
|
||||
_apply_toggle(session, parsed)
|
||||
|
||||
|
||||
async def _run_segment(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
message: SegmentPromptSource,
|
||||
) -> None:
|
||||
request = _build_generation_request(session, message, state)
|
||||
segment_idx = session.segment_idx
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentStart(
|
||||
segment_idx=segment_idx,
|
||||
prompt=message.prompt,
|
||||
total_steps=request.sampling.num_inference_steps,
|
||||
))
|
||||
|
||||
start = time.perf_counter()
|
||||
loop = asyncio.get_running_loop()
|
||||
# TODO: executor-wrapped generate() cannot be cancelled, so a
|
||||
# client disconnect mid-segment leaves the GPU work running to
|
||||
# completion. Real cancellation needs the generate_async API.
|
||||
try:
|
||||
result = await loop.run_in_executor(None, state.generator.generate, request)
|
||||
except Exception as exc:
|
||||
logger.exception("session %s: generator failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
elapsed_ms = (time.perf_counter() - start) * 1000.0
|
||||
|
||||
frames = _extract_frames(result)
|
||||
if not frames:
|
||||
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
|
||||
# Synchronous generator call has no per-step hook; emit one
|
||||
# terminal StepComplete so observability wiring still sees the
|
||||
# segment finish.
|
||||
total = request.sampling.num_inference_steps
|
||||
await _send_json(websocket, StepComplete(
|
||||
segment_idx=segment_idx,
|
||||
step=total,
|
||||
total_steps=total,
|
||||
stage="denoise",
|
||||
))
|
||||
|
||||
encoder = FragmentedMP4Encoder(
|
||||
width=request.sampling.width,
|
||||
height=request.sampling.height,
|
||||
fps=request.sampling.fps,
|
||||
segment_idx=segment_idx,
|
||||
)
|
||||
chunks_relayed = 0
|
||||
async with encoder:
|
||||
init_sent = False
|
||||
async for chunk in encoder.encode(frames):
|
||||
if chunk.kind == "init":
|
||||
await _send_json(websocket, MediaInit(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=chunk.stream_id,
|
||||
))
|
||||
init_sent = True
|
||||
await websocket.send_bytes(chunk.data)
|
||||
if init_sent and chunk.kind == "media":
|
||||
chunks_relayed += 1
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
MediaSegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=encoder.stream_id,
|
||||
chunks=chunks_relayed,
|
||||
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
|
||||
))
|
||||
|
||||
new_state = _extract_state(result)
|
||||
if new_state is not None:
|
||||
session.continuation_state = new_state
|
||||
state.session_store.store(session.id, new_state)
|
||||
|
||||
session.segment_idx += 1
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ACTIVE)
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
generation_time_ms=elapsed_ms,
|
||||
e2e_latency_ms=elapsed_ms,
|
||||
))
|
||||
|
||||
|
||||
def _build_stream_start(
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> Ltx2StreamStart:
|
||||
default = state.serve_config.default_request
|
||||
return Ltx2StreamStart(
|
||||
preset=session.preset,
|
||||
width=default.sampling.width,
|
||||
height=default.sampling.height,
|
||||
fps=default.sampling.fps,
|
||||
num_frames=default.sampling.num_frames,
|
||||
)
|
||||
|
||||
|
||||
def _build_generation_request(
|
||||
session: Session,
|
||||
message: SegmentPromptSource,
|
||||
state: ServerState,
|
||||
) -> GenerationRequest:
|
||||
# Start from the operator-pinned default_request to pick up the
|
||||
# preset-selected sampling knobs; override with per-message values.
|
||||
base = state.serve_config.default_request
|
||||
sampling_kwargs: dict[str, Any] = {
|
||||
"num_videos_per_prompt":
|
||||
base.sampling.num_videos_per_prompt,
|
||||
"seed":
|
||||
message.seed if message.seed is not None else base.sampling.seed,
|
||||
"num_frames":
|
||||
base.sampling.num_frames,
|
||||
"height":
|
||||
base.sampling.height,
|
||||
"width":
|
||||
base.sampling.width,
|
||||
"fps":
|
||||
base.sampling.fps,
|
||||
"num_inference_steps":
|
||||
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
|
||||
"guidance_scale":
|
||||
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
|
||||
}
|
||||
request = GenerationRequest(
|
||||
prompt=message.prompt,
|
||||
negative_prompt=message.negative_prompt or base.negative_prompt,
|
||||
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
|
||||
sampling=SamplingConfig(**sampling_kwargs),
|
||||
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
|
||||
state=session.continuation_state,
|
||||
)
|
||||
return request
|
||||
|
||||
|
||||
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
|
||||
kind = raw.get("kind")
|
||||
payload = raw.get("payload")
|
||||
if not isinstance(kind, str) or not isinstance(payload, dict):
|
||||
return None
|
||||
return ContinuationState(kind=kind, payload=payload)
|
||||
|
||||
|
||||
def _apply_toggle(session: Session, message: Any) -> None:
|
||||
if isinstance(message, EnhancementUpdated):
|
||||
session.enhancement_enabled = message.enabled
|
||||
elif isinstance(message, AutoExtensionUpdated):
|
||||
session.auto_extension_enabled = message.enabled
|
||||
elif isinstance(message, LoopGenerationUpdated):
|
||||
session.loop_generation_enabled = message.enabled
|
||||
elif isinstance(message, GenerationPausedUpdated):
|
||||
session.generation_paused = message.paused
|
||||
elif isinstance(message, SeedPromptsUpdated):
|
||||
session.curated_prompts = list(message.seed_prompts)
|
||||
|
||||
|
||||
def _extract_frames(result: Any) -> list:
|
||||
if hasattr(result, "frames"):
|
||||
return list(result.frames or [])
|
||||
if isinstance(result, dict):
|
||||
return list(result.get("frames") or [])
|
||||
return []
|
||||
|
||||
|
||||
def _extract_state(result: Any) -> ContinuationState | None:
|
||||
state = getattr(result, "state", None)
|
||||
if state is None and isinstance(result, dict):
|
||||
state = result.get("state")
|
||||
if isinstance(state, ContinuationState):
|
||||
return state
|
||||
if isinstance(state, dict):
|
||||
return _coerce_state(state)
|
||||
return None
|
||||
|
||||
|
||||
async def _send_json(websocket: WebSocket, message: Any) -> None:
|
||||
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
|
||||
await websocket.send_json(payload)
|
||||
|
||||
|
||||
async def _send_error(
|
||||
websocket: WebSocket,
|
||||
code: str,
|
||||
message: str,
|
||||
*,
|
||||
retryable: bool,
|
||||
) -> None:
|
||||
await _send_json(
|
||||
websocket,
|
||||
ErrorMessage(code=code, message=message, retryable=retryable),
|
||||
)
|
||||
|
||||
|
||||
def _cleanup_session(session: Session, state: ServerState) -> None:
|
||||
state.sessions.close(session.id)
|
||||
state.session_store.drop(session.id)
|
||||
init_image_path = session.metadata.get("session_init_image")
|
||||
if isinstance(init_image_path, str):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(init_image_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ServerState",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user