Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4a427692bf | ||
|
|
1214fb0f74 | ||
|
|
2dddbdf4c0 | ||
|
|
4aa065b96c | ||
|
|
7b1cd12059 | ||
|
|
d041b038bf | ||
|
|
9a15721d28 | ||
|
|
185d833d68 | ||
|
|
5568c90591 | ||
|
|
8031f27817 | ||
|
|
0b7e8b5d1d | ||
|
|
279e52ad8d | ||
|
|
6b3c1223c6 | ||
|
|
0c8687c919 | ||
|
|
ddbf41fa0e |
@@ -29,8 +29,8 @@
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
],
|
||||
"run_config": {
|
||||
"num_warmup_runs": 2,
|
||||
"num_measurement_runs": 5,
|
||||
"num_warmup_runs": 1,
|
||||
"num_measurement_runs": 3,
|
||||
"required_gpus": 2
|
||||
},
|
||||
"thresholds": {
|
||||
|
||||
@@ -63,72 +63,7 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
upload_performance_artifacts() {
|
||||
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
|
||||
LOCAL_DIR="downloaded_reports"
|
||||
|
||||
_download_reports() {
|
||||
log "Downloading perf_reports/ from Modal Volume..."
|
||||
mkdir -p "$LOCAL_DIR"
|
||||
if ! modal volume get hf-model-weights "perf_reports/" "$LOCAL_DIR"; then
|
||||
log "Error: Failed to download perf_reports/ from Modal Volume."
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_dashboard() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_*${SHORT_SHA}*" | head -n 1)
|
||||
log "TARGET dashboard: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
log "Found dashboard: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
buildkite-agent annotate --style info --context "perf-dashboard" < "$target"
|
||||
else
|
||||
log "Warning: Could not find a dashboard file matching $SHORT_SHA"
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_perf_summary() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "perf_*${SHORT_SHA}*" | head -n 1)
|
||||
log "TARGET perf summary: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
log "Found perf summary: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
buildkite-agent annotate --style info --context "perf-summary" < "$target"
|
||||
else
|
||||
log "Warning: Could not find a perf summary file matching $SHORT_SHA"
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_modal_volume() {
|
||||
log "Cleaning up perf_reports/ from Modal Volume..."
|
||||
if modal volume rm hf-model-weights "perf_reports/" --recursive; then
|
||||
log "Successfully deleted perf_reports/ from Modal Volume."
|
||||
else
|
||||
log "Warning: Failed to delete perf_reports/ from Modal Volume. Manual cleanup may be required."
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_local() {
|
||||
log "Cleaning up local download directory..."
|
||||
rm -rf "$LOCAL_DIR"
|
||||
}
|
||||
|
||||
# --- Main flow ---
|
||||
_download_reports || { _cleanup_local; return 1; }
|
||||
_upload_dashboard
|
||||
_upload_perf_summary
|
||||
_cleanup_modal_volume
|
||||
_cleanup_local
|
||||
}
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
@@ -189,9 +124,8 @@ case "$TEST_TYPE" in
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
|
||||
;;
|
||||
"performance")
|
||||
log "Running performance tests on Modal..."
|
||||
log "Running performance tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_performance_tests"
|
||||
POST_RUN_HOOK="upload_performance_artifacts"
|
||||
;;
|
||||
"api_server")
|
||||
log "Running API server integration tests..."
|
||||
@@ -213,10 +147,5 @@ else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
if [ -n "$POST_RUN_HOOK" ]; then
|
||||
log "Executing post-run hook: $POST_RUN_HOOK"
|
||||
"$POST_RUN_HOOK"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
|
||||
@@ -18,7 +18,6 @@ exclude: |
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
examples/.*|
|
||||
\.agents/.*|
|
||||
.github/workflows/publish-fastvideo.yml|
|
||||
.github/workflows/_template-build-image.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
|
||||
@@ -296,10 +296,8 @@ Action:
|
||||
|
||||
- Add or reuse a numerical parity test that loads the official model and the
|
||||
FastVideo model and compares outputs.
|
||||
- See examples in `tests/local_tests/` organized by model family
|
||||
(e.g., `tests/local_tests/sd35/`, `tests/local_tests/ltx2/`,
|
||||
`tests/local_tests/stable_audio/`) and the navigation index in
|
||||
`tests/local_tests/README.md`.
|
||||
- See examples in `tests/local_tests/` (e.g., `tests/local_tests/upsamplers/`)
|
||||
and the commands in `tests/local_tests/README.md`.
|
||||
- If there are discrepancies, add opt‑in logging to both models and compare
|
||||
activation summaries (layer output sums, per‑stage logs).
|
||||
- First align the loaded weights (validate `param_names_mapping`).
|
||||
@@ -350,8 +348,7 @@ Purpose:
|
||||
|
||||
Action:
|
||||
|
||||
- Add a pipeline parity test under `tests/local_tests/<family>/`
|
||||
(e.g., `tests/local_tests/<family>/test_<family>_pipeline_parity.py`).
|
||||
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
|
||||
- See the [Testing Guide](testing.md) for test conventions.
|
||||
|
||||
### 7) Add user‑facing examples
|
||||
|
||||
@@ -354,8 +354,6 @@ surfaces:
|
||||
return_frames: request.output.return_frames
|
||||
return_trajectory_latents: request.runtime.return_trajectory_latents
|
||||
return_trajectory_decoded: request.runtime.return_trajectory_decoded
|
||||
continuation_state: request.state
|
||||
return_continuation_state: request.output.return_state
|
||||
preset_owned:
|
||||
t_thresh: request.stage_overrides.refine.t_thresh
|
||||
spatial_refine_only: request.stage_overrides.refine.spatial_refine_only
|
||||
|
||||
@@ -1,177 +0,0 @@
|
||||
# Streaming WebSocket Server Contract
|
||||
|
||||
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
|
||||
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
|
||||
document is the authoritative spec for the message catalogue and the
|
||||
session state machine. Any change to either must update this document
|
||||
in the same PR that touches `protocol.py` or `session.py`.
|
||||
|
||||
## Endpoint
|
||||
|
||||
| Path | Protocol | Purpose |
|
||||
|---|---|---|
|
||||
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
|
||||
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
|
||||
|
||||
The server is launched by `fastvideo serve --config <serve.yaml>` when
|
||||
the config carries a `streaming:` block. Without that block the same CLI
|
||||
launches the OpenAI stateless HTTP server instead.
|
||||
|
||||
## Connection lifecycle
|
||||
|
||||
Every WebSocket connection holds exactly one `Session`. Sessions move
|
||||
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
|
||||
|
||||
```
|
||||
┌──────────────┐
|
||||
│ INITIALIZING │ ← WebSocket accepted, before init frame
|
||||
└──────┬───────┘
|
||||
│ session_init_v2 received
|
||||
┌──────────────┼──────────────┐
|
||||
▼ ▼ ▼
|
||||
QUEUED GPU_BINDING REJECTED
|
||||
│ │ ↑
|
||||
│ slot ready │ │ max-sessions hit
|
||||
▼ ▼ │ or invalid init
|
||||
┌────────┐ │
|
||||
│ ACTIVE │ ────────┘
|
||||
└────┬───┘
|
||||
segment loop │
|
||||
│
|
||||
┌───────────┼───────────┐
|
||||
▼ ▼ ▼
|
||||
COMPLETE ERROR TIMEOUT
|
||||
(clean leave) (any failure) (idle / segment_cap reached)
|
||||
```
|
||||
|
||||
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
|
||||
no transitions out. The transition matrix is enforced in
|
||||
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
|
||||
|
||||
`SessionManager` enforces the per-process budgets pulled from
|
||||
`StreamingConfig`:
|
||||
|
||||
- `session_timeout_seconds` — idle reaper drops sessions that haven't
|
||||
advanced; non-terminal sessions transition to `TIMEOUT`.
|
||||
- `generation_segment_cap` — a session that hits the cap transitions to
|
||||
`COMPLETE` after the last segment ships.
|
||||
|
||||
## Message catalogue
|
||||
|
||||
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
|
||||
`protocol.py` are the source of truth; this table is the human-readable
|
||||
view.
|
||||
|
||||
### Client → server
|
||||
|
||||
| `type` | Required fields | Purpose |
|
||||
|---|---|---|
|
||||
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
|
||||
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
|
||||
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
|
||||
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
|
||||
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
|
||||
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
|
||||
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
|
||||
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
|
||||
|
||||
The opening frame must be `session_init_v2`. Any other first frame is
|
||||
rejected with an `error` (code `invalid_message`) and the WebSocket is
|
||||
closed.
|
||||
|
||||
### Server → client
|
||||
|
||||
| `type` | Carries | When emitted |
|
||||
|---|---|---|
|
||||
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
|
||||
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
|
||||
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
|
||||
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
|
||||
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
|
||||
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
|
||||
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
|
||||
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
|
||||
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
|
||||
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
|
||||
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
|
||||
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
|
||||
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
|
||||
|
||||
## Continuation state
|
||||
|
||||
The session optionally accepts a `continuation_state` dict inside the
|
||||
opening `session_init_v2` frame. When present, the server hydrates it
|
||||
into a `ContinuationState(kind, payload)` envelope and feeds it as the
|
||||
`request.state` on the first segment's `GenerationRequest` — letting a
|
||||
client resume after a disconnect, migrate sessions across processes,
|
||||
or replay a prior session.
|
||||
|
||||
After every segment, if the runtime returns a fresh state, the server
|
||||
persists it to the `SessionStore` so a `snapshot_state` request can
|
||||
export it. The store and serialization contracts live with the model
|
||||
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
|
||||
|
||||
## Example flow
|
||||
|
||||
```
|
||||
client server
|
||||
────── ──────
|
||||
WS /v1/stream ─────── connect ─────────────────────────►
|
||||
◄────── (accept)
|
||||
|
||||
{"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox in snow", "the fox jumps"],
|
||||
"initial_image": {...},
|
||||
"stream_mode": "av_fmp4"} ─────────────────────────────►
|
||||
|
||||
(validate, queue, bind)
|
||||
◄──── {"type": "queue_status",
|
||||
"position": 0, "queue_depth": 0}
|
||||
◄──── {"type": "gpu_assigned",
|
||||
"gpu_id": 0, "model_id": "..."}
|
||||
◄──── {"type": "ltx2_stream_start", ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "a fox in snow",
|
||||
"source": "curated"} ───────────────────────────────────►
|
||||
(run pipeline)
|
||||
◄──── {"type": "ltx2_segment_start",
|
||||
"segment_idx": 1, ...}
|
||||
◄──── {"type": "step_complete",
|
||||
"segment_idx": 1, "timings": {...}}
|
||||
◄──── {"type": "media_init",
|
||||
"segment_idx": 1,
|
||||
"mime": "video/mp4", ...}
|
||||
◄──── <binary fMP4 init segment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── {"type": "media_segment_complete",
|
||||
"segment_idx": 1, "chunks": 12}
|
||||
◄──── {"type": "ltx2_segment_complete",
|
||||
"segment_idx": 1, ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "the fox jumps"} ─────────────────────────────►
|
||||
(segment 2 …)
|
||||
|
||||
{"type": "snapshot_state"} ──────────────────────────────►
|
||||
◄──── {"type": "continuation_state_snapshot",
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, ...}}
|
||||
|
||||
(close) ──────────────────────────────────────────────────►
|
||||
(session → COMPLETE)
|
||||
```
|
||||
|
||||
## Backward / forward compatibility
|
||||
|
||||
- Adding a new client message: append a Pydantic model to `protocol.py`
|
||||
with a unique `type`; add the discriminator entry to `ClientMessage`;
|
||||
add a row to the table above. Old clients that don't send the new
|
||||
message remain compatible.
|
||||
- Adding a new server message: emit only when a new feature flag is
|
||||
enabled (or always emit, since clients ignore unknown types).
|
||||
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
|
||||
→ `session_init_v3`) and accept both for one release cycle. Never
|
||||
silently change field semantics under the same `type`.
|
||||
@@ -86,25 +86,3 @@ sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
|
||||
- Learning rate: 2e-5
|
||||
- Training steps: 3000 (~12 hours)
|
||||
- HSDP shard dim: 1
|
||||
|
||||
## 🧭 Note on `real_score_guidance_scale`
|
||||
|
||||
The teacher CFG used inside the DMD loss follows the DMD2 reference
|
||||
implementation and uses the parameterization
|
||||
|
||||
```
|
||||
x = x_cond + w * (x_cond - x_uncond)
|
||||
```
|
||||
|
||||
rather than the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)`. The
|
||||
two are mathematically equivalent up to a constant offset:
|
||||
|
||||
| `real_score_guidance_scale` (`w`) | Equivalent standard CFG (`w + 1`) | Output |
|
||||
|-----------------------------------|-----------------------------------|-----------------------|
|
||||
| `-1` | `0` | unconditional |
|
||||
| `0` | `1` | conditional |
|
||||
| `3.5` (default) | `4.5` | strong guidance |
|
||||
|
||||
So `real_score_guidance_scale` should be read as the **extra** guidance
|
||||
strength added on top of the conditional prediction. When porting values
|
||||
from a paper that uses the Ho & Salimans form, subtract 1.
|
||||
|
||||
@@ -1,77 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — text-to-audio (baseline) example.
|
||||
|
||||
User story (game-audio designer, prototyping):
|
||||
"I'm prototyping a level and I need 6 seconds of background
|
||||
ambience — gentle wind, distant thunder, a hint of birdsong. I
|
||||
don't want to dig through a sound library; I want to type what I
|
||||
hear in my head and get a wav back. If it's wrong I'll iterate
|
||||
on the prompt. This is the first stop."
|
||||
|
||||
User story (musician sketching ideas):
|
||||
"I want to bounce a 30s lo-fi drum loop to use as a placeholder
|
||||
bed while I build the rest of the track. Type prompt, get audio,
|
||||
drop into the DAW. The actual production beat I'll record
|
||||
myself, but I need *something* to write the chords against."
|
||||
|
||||
User story (researcher exploring the model):
|
||||
"First time touching Stable Audio Open — what does it sound
|
||||
like at default settings? This is the smallest amount of code
|
||||
that goes from prompt to mp4."
|
||||
|
||||
How it works:
|
||||
Pure text-to-audio (T2A). The pipeline runs:
|
||||
T5 + NumberConditioner -> StableAudioDiT -> Oobleck VAE
|
||||
via the `dpmpp-3m-sde` k-diffusion sampler. All components are
|
||||
FastVideo-native — no diffusers / transformers model imports at
|
||||
runtime (see REVIEW item 30). Mirrors upstream
|
||||
`stable_audio_tools.inference.generation.generate_diffusion_cond`
|
||||
bit-for-bit (~0.2% abs_mean drift on 25 steps).
|
||||
|
||||
Tunable knobs (the "creative dials"):
|
||||
audio_end_in_s
|
||||
1–6 — quick ideation (sub-10s wall clock at 100 steps)
|
||||
10–30 — full musical phrase / loop length (the README example
|
||||
uses 30s)
|
||||
47.5 — model maximum (full sample_size = 2097152 / 44100 Hz)
|
||||
num_inference_steps
|
||||
25 — fast preview, occasional artifacts
|
||||
100 — preset default (matches the HF model card)
|
||||
250 — diminishing returns past here
|
||||
guidance_scale
|
||||
3 — looser, more variation per seed
|
||||
7 — preset default; matches README
|
||||
12+ — sharper but can sound "fried"
|
||||
|
||||
Prerequisites:
|
||||
1. Accept the terms on https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
and export your HF token in the shell:
|
||||
export HF_TOKEN=hf_...
|
||||
2. Install optional inference deps (one-time):
|
||||
pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_audio/stable_audio_basic/output_stable_audio.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# 6-second clip; the model max is ~47.5s.
|
||||
audio_end_in_s=6.0,
|
||||
# The registered preset gives 100 steps + CFG=7.0 by default;
|
||||
# override num_inference_steps / guidance_scale here for QA.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,77 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — audio-to-audio variation example.
|
||||
|
||||
User story (musician, late at night):
|
||||
"I generated this 12-second lo-fi loop earlier and I love the chord
|
||||
progression and overall vibe, but the snare hit at 0:08 sounds wrong
|
||||
and the rhythm feels stiff. I don't want to start over from scratch
|
||||
and lose what's working — I want the model to keep the harmony and
|
||||
mood but reroll the percussion + groove."
|
||||
|
||||
User story (sound designer, on a deadline):
|
||||
"I have one good 'sword clang' SFX. The art director wants 8 sibling
|
||||
variations that all feel like the same sword from different angles —
|
||||
same metal, same weight, slightly different impact. I'd rather
|
||||
refine my one good take than text-prompt my way through 50 misses."
|
||||
|
||||
Pass `init_audio=path/to/clip` (any wav/mp3/mp4/m4a/flac the standard
|
||||
deps decode) and the model will use it as a starting point for the
|
||||
text prompt instead of pure noise.
|
||||
|
||||
Picking `init_audio_strength` (0.0 to 1.0):
|
||||
|
||||
Higher = closer to the source clip. Lower = more transformation.
|
||||
(Same convention as the "Input Audio Strength" slider in
|
||||
Stability's commercial Stable Audio web UI, so values transfer
|
||||
directly.)
|
||||
|
||||
| strength | what you get |
|
||||
|----------|----------------------------------------------------|
|
||||
| 1.00 | Output ≈ reference. No transformation. |
|
||||
| 0.85 | Texture micro-variation only. |
|
||||
| 0.70 | Light reroll, same instruments. |
|
||||
| 0.60 | Default. Instrument identity is replaceable |
|
||||
| | (cello can take over from piano on the same notes).|
|
||||
| 0.50 | Heavy — only melody / chord progression survives. |
|
||||
| 0.30 | Reference acts as a loose mood prompt. |
|
||||
| 0.00 | Plain T2A — reference ignored. |
|
||||
|
||||
Rule of thumb by intent:
|
||||
* "Fix one part of this clip" -> 0.75 .. 0.85
|
||||
* "Same notes, different instrument" -> 0.55 .. 0.65
|
||||
* "Same chord progression, new content" -> 0.40 .. 0.55
|
||||
* "Use this as a loose mood prompt" -> 0.20 .. 0.35
|
||||
|
||||
If the reference timbre is bleeding through more than you want,
|
||||
lower it; if the structure is gone, raise it.
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Change the piano to a cello playing the same notes"
|
||||
# Path to any audio-bearing file (wav, mp3, mp4, m4a, flac, ...).
|
||||
# Set to `None` to skip A2A and run plain T2A.
|
||||
INIT_AUDIO_PATH: str | None = None
|
||||
# Reference fidelity in [0, 1] -- higher = closer to source.
|
||||
INIT_AUDIO_STRENGTH = 0.6
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_a2a/output_a2a.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=6.0,
|
||||
init_audio=INIT_AUDIO_PATH,
|
||||
init_audio_strength=INIT_AUDIO_STRENGTH,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,84 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — inpainting / outpainting (loop extension) example.
|
||||
|
||||
User story (loop extension — the killer app):
|
||||
"I have a 6-second drum loop my client likes. They want it as
|
||||
background bed for a 30-second ad. I need it to loop seamlessly,
|
||||
but a hard cut every 6s sounds bad. Let me extend it to 30s,
|
||||
keeping the first 6s exactly as-is and letting the model continue
|
||||
the groove for the remaining 24s."
|
||||
|
||||
User story (audio repair):
|
||||
"There's a microphone bump at 0:14 in this 30-second field
|
||||
recording — really obvious in headphones. Mask out 0:13 to 0:15
|
||||
and let the model regenerate plausible ambience that blends in.
|
||||
Everything else stays exactly as I recorded it."
|
||||
|
||||
User story (transition smoothing):
|
||||
"I have two 10-second clips I want to crossfade. Mask out a 1s
|
||||
overlap region in the middle and let the model invent a coherent
|
||||
transition between the two."
|
||||
|
||||
How it works (RePaint-style blending):
|
||||
Stable Audio Open 1.0 wasn't trained as an inpainting model
|
||||
(`model_type=diffusion_cond`, not `diffusion_cond_inpaint`), so we
|
||||
can't use the upstream's mask-conditioned approach directly. We
|
||||
use the RePaint trick instead, which works on any v-prediction
|
||||
diffusion model:
|
||||
|
||||
1. Encode the reference clip into latent space.
|
||||
2. At every denoising step `i`, replace the kept region of the
|
||||
in-flight latent (where mask == 1) with the reference
|
||||
re-noised to the next timestep's sigma. Only the unkept
|
||||
region (mask == 0) is freely denoised.
|
||||
3. After the loop, the kept region is exactly the reference;
|
||||
the unkept region is freshly generated content.
|
||||
|
||||
This is approximate compared to a properly trained inpainting
|
||||
checkpoint — the seam between kept/unkept can have slight EQ
|
||||
discontinuity — but it works on the existing public model.
|
||||
|
||||
Tunable: the mask is a 1-D tensor in {0, 1} at the model's sample
|
||||
rate. Conventions:
|
||||
1.0 = keep this sample from the reference
|
||||
0.0 = regenerate this sample
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
|
||||
# Required: path to the reference audio file (wav, mp3, mp4, m4a, flac,
|
||||
# ...) you want to extend or repair. The pipeline raises if a mask is
|
||||
# passed without a reference, so this must be a real path.
|
||||
REFERENCE_AUDIO_PATH = "path/to/your/loop.wav"
|
||||
KEEP_SECONDS = 6.0 # first KEEP_SECONDS preserved exactly
|
||||
TOTAL_SECONDS = 12.0 # extend the loop to this duration
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not os.path.isfile(REFERENCE_AUDIO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
|
||||
"Edit this script to point at a real audio file (wav/mp3/mp4/"
|
||||
"m4a/flac) before running.")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_inpaint/output_inpaint.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=TOTAL_SECONDS,
|
||||
inpaint_audio=REFERENCE_AUDIO_PATH,
|
||||
# Tuple form: keep first KEEP_SECONDS, regenerate the rest.
|
||||
inpaint_mask=(KEEP_SECONDS, TOTAL_SECONDS),
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,53 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open Small — fast / lightweight T2A example.
|
||||
|
||||
User story (interactive UI builder):
|
||||
"I'm building a sound-design UI where the user types a prompt and
|
||||
we want sub-2-second feedback so the experience feels like
|
||||
autocomplete, not a render queue. The full Stable Audio Open 1.0
|
||||
takes ~8s on a single GPU; the small variant takes a fraction of
|
||||
that — quality is lower but completely usable for real-time
|
||||
iteration."
|
||||
|
||||
User story (overnight batch jobs):
|
||||
"I'm generating 10,000 short SFX variants for a procedural game.
|
||||
Wall-clock matters more than per-clip polish — give me the small
|
||||
model so I can fit the run in one night instead of a week."
|
||||
|
||||
How it works:
|
||||
The small variant is a separate Stability AI checkpoint
|
||||
(`stabilityai/stable-audio-open-small`) that ships the same Oobleck
|
||||
VAE as the 1.0 base model but a smaller / faster DiT (`embed_dim=1024`,
|
||||
`depth=16`, `qk_norm="ln"`) and only one duration conditioner
|
||||
(`seconds_total`, no `seconds_start`). FastVideo loads from the
|
||||
converted Diffusers-format repo `FastVideo/stable-audio-open-small-Diffusers`
|
||||
via the standard component loader; per-variant arch fields come
|
||||
from `transformer/config.json` and `conditioner/config.json`.
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`. The converted repo is
|
||||
public so no gated-access flow is required.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-small-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_audio/stable_audio_small/output_stable_audio_small.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Small variant trains on a ~11.9s window — keep `audio_end_in_s`
|
||||
# at or below that.
|
||||
audio_end_in_s=6.0,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,73 +0,0 @@
|
||||
# Cosmos Predict2 2B T2V finetune config.
|
||||
#
|
||||
# Data must be preprocessed with Cosmos VAE + T5 text encoder
|
||||
# into parquet format before training.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: nvidia/Cosmos-Predict2-2B-Video2World
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos_preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
# Cosmos VAE: 4x temporal, 8x spatial compression.
|
||||
# 93 frames -> 24 latent frames, 480x832 -> 60x104
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 5000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos_finetune
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos
|
||||
run_name: cosmos_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos_pipeline.Cosmos2VideoToWorldPipeline
|
||||
dataset_file: data/cosmos_preprocessed/validation_prompts.json
|
||||
every_steps: 100
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -1,79 +0,0 @@
|
||||
# Cosmos-Predict2.5-2B Text-to-World overfitting test config.
|
||||
#
|
||||
# Overfits on a few short videos (480x832, 93 frames) to verify the
|
||||
# Cosmos 2.5 training plugin works end-to-end.
|
||||
#
|
||||
# Preprocess data first:
|
||||
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# bash examples/train/run.sh examples/train/configs/overfit_cosmos25_t2w.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: KyleShao/Cosmos-Predict2.5-2B-Diffusers
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
flow_shift: 1.0
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos25_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos25_overfit
|
||||
training_state_checkpointing_steps: 50
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos25
|
||||
run_name: cosmos25_overfit
|
||||
|
||||
model:
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos2_5_pipeline.Cosmos2_5Pipeline
|
||||
dataset_file: data/cosmos25_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 150
|
||||
sampling_steps: [35]
|
||||
guidance_scale: 7.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -5,11 +5,6 @@ from fastvideo_kernel.ops import (
|
||||
video_sparse_attn,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn import (
|
||||
block_sparse_attn,
|
||||
block_sparse_attn_from_indices,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.vmoba import (
|
||||
moba_attn_varlen,
|
||||
process_moba_input,
|
||||
@@ -27,8 +22,6 @@ from fastvideo_kernel.turbodiffusion_ops import (
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"block_sparse_attn",
|
||||
"block_sparse_attn_from_indices",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
@@ -8,11 +6,6 @@ from typing import Tuple
|
||||
import torch
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Backend selection helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_sm90_ops():
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
|
||||
@@ -32,66 +25,38 @@ def _is_sm90() -> bool:
|
||||
|
||||
|
||||
def _force_triton() -> bool:
|
||||
# Force Triton even on SM90 and even if the compiled extension is available.
|
||||
# Useful for CI / debugging / parity testing.
|
||||
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Index helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Compact a bool block_map to (q2k_idx, q2k_num). Legacy path only."""
|
||||
"""
|
||||
Preferred map->index conversion used by the wrapper.
|
||||
|
||||
This wrapper **requires** the Triton implementation.
|
||||
If Triton (or the Triton map_to_index module) is not available, it raises.
|
||||
"""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
if block_map.dim() != 4:
|
||||
raise ValueError(
|
||||
f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), "
|
||||
f"got shape={tuple(block_map.shape)}"
|
||||
)
|
||||
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
|
||||
if block_map.dtype != torch.bool:
|
||||
block_map = block_map.to(torch.bool)
|
||||
|
||||
if not block_map.is_cuda:
|
||||
raise RuntimeError(
|
||||
"block_map must be a CUDA tensor (Triton map_to_index required)."
|
||||
)
|
||||
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
|
||||
|
||||
try:
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index
|
||||
except Exception as e: # pragma: no cover - environment issue
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
|
||||
except Exception as e:
|
||||
raise ImportError(
|
||||
"Triton map_to_index is required but not available. "
|
||||
"Ensure Triton is installed and "
|
||||
"fastvideo_kernel.triton_kernels.index is importable."
|
||||
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
|
||||
) from e
|
||||
return triton_map_to_index(block_map)
|
||||
|
||||
|
||||
def _invert_indices_for_backward(
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from fastvideo_kernel.triton_kernels.index import invert_indices
|
||||
return invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
|
||||
|
||||
|
||||
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
|
||||
"""Return `t` as a contiguous int32 tensor, raising a clear error on CPU input."""
|
||||
if not t.is_cuda:
|
||||
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
|
||||
if t.dtype != torch.int32:
|
||||
t = t.to(torch.int32)
|
||||
if not t.is_contiguous():
|
||||
t = t.contiguous()
|
||||
return t
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Triton backend custom ops (index-native)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo_kernel::block_sparse_attn_triton",
|
||||
mutates_args=(),
|
||||
@@ -101,40 +66,34 @@ def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
triton_block_sparse_attn_forward,
|
||||
)
|
||||
|
||||
o, M = triton_block_sparse_attn_forward(
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
variable_block_sizes,
|
||||
)
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty(
|
||||
(q.shape[0], q.shape[1], q.shape[2]),
|
||||
device=q.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@@ -150,32 +109,20 @@ def block_sparse_attn_backward_triton(
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
|
||||
grad_output = grad_output.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
triton_block_sparse_attn_backward,
|
||||
)
|
||||
|
||||
num_kv_blocks = int(variable_block_sizes.numel())
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(
|
||||
q2k_idx, q2k_num, num_kv_blocks
|
||||
)
|
||||
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
|
||||
# o/M are kernel outputs so are already contiguous.
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(
|
||||
grad_output.contiguous(),
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
o,
|
||||
M,
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes,
|
||||
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv
|
||||
|
||||
@@ -188,8 +135,7 @@ def _block_sparse_attn_backward_triton_fake(
|
||||
v: torch.Tensor,
|
||||
o: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q)
|
||||
@@ -198,28 +144,19 @@ def _block_sparse_attn_backward_triton_fake(
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _setup_context_triton(ctx, inputs, output):
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
|
||||
o, M = output
|
||||
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
def _backward_triton(ctx, grad_o, grad_M):
|
||||
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(
|
||||
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None, None
|
||||
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
block_sparse_attn_triton.register_autograd(
|
||||
_backward_triton, setup_context=_setup_context_triton
|
||||
)
|
||||
def _setup_context_triton(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
o, M = output
|
||||
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SM90 backend custom ops (index-native)
|
||||
# ---------------------------------------------------------------------------
|
||||
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
@@ -231,21 +168,21 @@ def block_sparse_attn_sm90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
block_sparse_fwd, _ = _get_sm90_ops()
|
||||
if block_sparse_fwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
|
||||
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
o_padded, lse_padded = block_sparse_fwd(
|
||||
q_padded.contiguous(),
|
||||
k_padded.contiguous(),
|
||||
v_padded.contiguous(),
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
variable_block_sizes,
|
||||
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
|
||||
)
|
||||
return o_padded, lse_padded
|
||||
|
||||
@@ -255,16 +192,11 @@ def _block_sparse_attn_sm90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
o = torch.empty_like(q_padded)
|
||||
lse = torch.empty(
|
||||
(q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1),
|
||||
device=q_padded.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
|
||||
return o, lse
|
||||
|
||||
|
||||
@@ -280,34 +212,30 @@ def block_sparse_attn_backward_sm90(
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
_, block_sparse_bwd = _get_sm90_ops()
|
||||
if block_sparse_bwd is None:
|
||||
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
|
||||
|
||||
num_kv_blocks = int(variable_block_sizes.numel())
|
||||
k2q_idx, k2q_num = _invert_indices_for_backward(
|
||||
q2k_idx, q2k_num, num_kv_blocks
|
||||
)
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
block_map = block_map.to(torch.bool)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
|
||||
dq, dk, dv = block_sparse_bwd(
|
||||
q_padded.contiguous(),
|
||||
k_padded.contiguous(),
|
||||
v_padded.contiguous(),
|
||||
q_padded,
|
||||
k_padded,
|
||||
v_padded,
|
||||
o_padded,
|
||||
lse_padded,
|
||||
grad_output_padded.contiguous(),
|
||||
grad_output_padded,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
variable_block_sizes,
|
||||
variable_block_sizes.int(),
|
||||
)
|
||||
# C++ kernel returns fp32 grads; cast back to the input dtype.
|
||||
out_dtype = grad_output_padded.dtype
|
||||
return dq.to(out_dtype), dk.to(out_dtype), dv.to(out_dtype)
|
||||
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
|
||||
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
|
||||
|
||||
|
||||
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
|
||||
@@ -318,8 +246,7 @@ def _block_sparse_attn_backward_sm90_fake(
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
dq = torch.empty_like(q_padded)
|
||||
@@ -328,57 +255,21 @@ def _block_sparse_attn_backward_sm90_fake(
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def _setup_context_sm90(ctx, inputs, output):
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
|
||||
o, lse = output
|
||||
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
|
||||
|
||||
|
||||
def _backward_sm90(ctx, grad_o, grad_lse):
|
||||
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
|
||||
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_sm90(
|
||||
grad_o, q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes
|
||||
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
|
||||
)
|
||||
return dq, dk, dv, None, None, None
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
block_sparse_attn_sm90.register_autograd(
|
||||
_backward_sm90, setup_context=_setup_context_sm90
|
||||
)
|
||||
def _setup_context_sm90(ctx, inputs, output):
|
||||
q, k, v, block_map, variable_block_sizes = inputs
|
||||
o, lse = output
|
||||
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def block_sparse_attn_from_indices(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Block-sparse attention with autograd, taking compact per-row KV indices."""
|
||||
# Normalize index tensors once at the public boundary so the custom ops
|
||||
# and their fakes can assume int32/contiguous. No-op on well-formed input.
|
||||
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
|
||||
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
|
||||
variable_block_sizes = _as_int32_contig(variable_block_sizes, "variable_block_sizes")
|
||||
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
use_sm90 = (
|
||||
(not _force_triton())
|
||||
and _is_sm90()
|
||||
and block_sparse_fwd is not None
|
||||
and block_sparse_bwd is not None
|
||||
)
|
||||
if use_sm90:
|
||||
return block_sparse_attn_sm90(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
|
||||
# to a multiple of the block size (64 tokens).
|
||||
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
|
||||
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
|
||||
|
||||
|
||||
def block_sparse_attn(
|
||||
@@ -388,8 +279,16 @@ def block_sparse_attn(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Bool-mask compat wrapper; prefer block_sparse_attn_from_indices."""
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
return block_sparse_attn_from_indices(
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes
|
||||
)
|
||||
"""
|
||||
Unified block-sparse attention op with autograd support.
|
||||
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
|
||||
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
|
||||
"""
|
||||
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
|
||||
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
|
||||
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
|
||||
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
|
||||
# to a multiple of the block size (64 tokens).
|
||||
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import math
|
||||
import torch
|
||||
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
|
||||
from .block_sparse_attn import block_sparse_attn
|
||||
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
|
||||
|
||||
# Try to load the C++ extension
|
||||
@@ -125,18 +125,13 @@ def video_sparse_attn(
|
||||
out_c = out_c.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch, heads, q_seq_len, dim)
|
||||
|
||||
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
|
||||
# Sparse branch
|
||||
topk_idx = torch.topk(scores, topk, dim=-1).indices
|
||||
q2k_idx = topk_idx.to(torch.int32).contiguous()
|
||||
q2k_num = torch.full(
|
||||
(batch, heads, q_num_blocks),
|
||||
topk,
|
||||
dtype=torch.int32,
|
||||
device=q.device,
|
||||
)
|
||||
out_s = block_sparse_attn_from_indices(
|
||||
q, k, v, q2k_idx, q2k_num, variable_block_sizes
|
||||
)[0]
|
||||
mask = torch.zeros_like(scores,
|
||||
dtype=torch.bool).scatter_(-1, topk_idx, True)
|
||||
|
||||
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
## pytorch sdpa version of block sparse ##
|
||||
from typing import Tuple
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
@@ -154,114 +153,3 @@ def map_to_index(block_map: torch.Tensor):
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _invert_indices_kernel(
|
||||
q2k_idx_ptr,
|
||||
q2k_num_ptr,
|
||||
k2q_idx_ptr,
|
||||
k2q_num_ptr,
|
||||
q2k_idx_b, q2k_idx_h, q2k_idx_q, q2k_idx_k,
|
||||
q2k_num_b, q2k_num_h, q2k_num_q,
|
||||
k2q_idx_b, k2q_idx_h, k2q_idx_k, k2q_idx_q,
|
||||
k2q_num_b, k2q_num_h, k2q_num_k,
|
||||
MAX_KV_PER_Q: tl.constexpr,
|
||||
):
|
||||
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
|
||||
pid_b = tl.program_id(0)
|
||||
pid_h = tl.program_id(1)
|
||||
pid_q = tl.program_id(2)
|
||||
|
||||
n = tl.load(
|
||||
q2k_num_ptr
|
||||
+ pid_b * q2k_num_b
|
||||
+ pid_h * q2k_num_h
|
||||
+ pid_q * q2k_num_q
|
||||
)
|
||||
|
||||
q2k_row = (
|
||||
q2k_idx_ptr
|
||||
+ pid_b * q2k_idx_b
|
||||
+ pid_h * q2k_idx_h
|
||||
+ pid_q * q2k_idx_q
|
||||
)
|
||||
|
||||
for i in tl.range(0, MAX_KV_PER_Q):
|
||||
if i < n:
|
||||
kv = tl.load(q2k_row + i * q2k_idx_k)
|
||||
count_ptr = (
|
||||
k2q_num_ptr
|
||||
+ pid_b * k2q_num_b
|
||||
+ pid_h * k2q_num_h
|
||||
+ kv * k2q_num_k
|
||||
)
|
||||
pos = tl.atomic_add(count_ptr, 1)
|
||||
tl.store(
|
||||
k2q_idx_ptr
|
||||
+ pid_b * k2q_idx_b
|
||||
+ pid_h * k2q_idx_h
|
||||
+ kv * k2q_idx_k
|
||||
+ pos * k2q_idx_q,
|
||||
pid_q,
|
||||
)
|
||||
|
||||
|
||||
def invert_indices(
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
|
||||
if q2k_idx.dim() != 4:
|
||||
raise ValueError(
|
||||
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
|
||||
)
|
||||
if q2k_num.dim() != 3:
|
||||
raise ValueError(
|
||||
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
|
||||
)
|
||||
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
|
||||
raise RuntimeError("invert_indices requires CUDA tensors.")
|
||||
|
||||
B, H, Nq, Mk = q2k_idx.shape
|
||||
if q2k_num.shape != (B, H, Nq):
|
||||
raise ValueError(
|
||||
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
|
||||
f"[B, H, Nq] = {(B, H, Nq)}"
|
||||
)
|
||||
|
||||
q2k_idx = q2k_idx.contiguous()
|
||||
q2k_num = q2k_num.contiguous()
|
||||
if q2k_idx.dtype != torch.int32:
|
||||
q2k_idx = q2k_idx.to(torch.int32)
|
||||
if q2k_num.dtype != torch.int32:
|
||||
q2k_num = q2k_num.to(torch.int32)
|
||||
|
||||
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
|
||||
# `Nq` is a tight upper bound on the compacted K->Q slots.
|
||||
k2q_idx = torch.empty(
|
||||
(B, H, num_kv_blocks, Nq),
|
||||
dtype=torch.int32,
|
||||
device=q2k_idx.device,
|
||||
)
|
||||
k2q_num = torch.zeros(
|
||||
(B, H, num_kv_blocks),
|
||||
dtype=torch.int32,
|
||||
device=q2k_idx.device,
|
||||
)
|
||||
|
||||
grid = (B, H, Nq)
|
||||
_invert_indices_kernel[grid](
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
k2q_idx,
|
||||
k2q_num,
|
||||
q2k_idx.stride(0), q2k_idx.stride(1), q2k_idx.stride(2), q2k_idx.stride(3),
|
||||
q2k_num.stride(0), q2k_num.stride(1), q2k_num.stride(2),
|
||||
k2q_idx.stride(0), k2q_idx.stride(1), k2q_idx.stride(2), k2q_idx.stride(3),
|
||||
k2q_num.stride(0), k2q_num.stride(1), k2q_num.stride(2),
|
||||
MAX_KV_PER_Q=Mk,
|
||||
)
|
||||
|
||||
return k2q_idx, k2q_num
|
||||
|
||||
+8
-63
@@ -17,7 +17,6 @@ from fastvideo.api.request_metadata import (
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
CompileConfig,
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
@@ -27,10 +26,7 @@ from fastvideo.api.schema import (
|
||||
)
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import REFINE_FLAT_KEYS
|
||||
from fastvideo.utils import shallow_asdict
|
||||
|
||||
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
|
||||
@@ -44,10 +40,7 @@ _LEGACY_REQUEST_ALIASES = {
|
||||
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
|
||||
"embedded_cfg_scale",
|
||||
})
|
||||
# torch.compile kwargs that map to first-class CompileConfig fields.
|
||||
_COMPILE_TYPED_KEYS = ("backend", "fullgraph", "mode", "dynamic")
|
||||
# LTX-2 refine flat kwargs (init + per-request) known to FastVideoArgs.
|
||||
_LTX2_REFINE_FLAT_KEYS = (refine_preset_override_fields() | refine_stage_override_fields())
|
||||
|
||||
|
||||
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
|
||||
@@ -118,10 +111,8 @@ def legacy_from_pretrained_to_config(
|
||||
offload["pin_cpu_memory"] = value
|
||||
elif key == "enable_torch_compile":
|
||||
compile_config["enabled"] = value
|
||||
elif key == "enable_torch_compile_text_encoder":
|
||||
compile_config["text_encoder_enabled"] = value
|
||||
elif key == "torch_compile_kwargs":
|
||||
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
|
||||
remaining: dict[str, Any] = (dict(value) if isinstance(value, Mapping) else {})
|
||||
for first_class in _COMPILE_TYPED_KEYS:
|
||||
if first_class in remaining:
|
||||
compile_config[first_class] = remaining.pop(first_class)
|
||||
@@ -237,12 +228,6 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
kwargs["workload_type"] = normalized.pipeline.workload_type
|
||||
if normalized.pipeline.vae_tiling is not None:
|
||||
kwargs["ltx2_vae_tiling"] = normalized.pipeline.vae_tiling
|
||||
if engine.compile.text_encoder_enabled is not None:
|
||||
# ``FastVideoArgs.from_kwargs`` filters to declared fields, so
|
||||
# this is a no-op on the current legacy path. Emit anyway so the
|
||||
# realtime runtime (PR 7.6) — which reads from the kwargs dict
|
||||
# before FastVideoArgs filtering — can pick it up once wired.
|
||||
kwargs["enable_torch_compile_text_encoder"] = (engine.compile.text_encoder_enabled)
|
||||
|
||||
quantization = engine.quantization
|
||||
if quantization is not None and quantization.text_encoder_quant is not None:
|
||||
@@ -273,7 +258,7 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
|
||||
refine = preset_overrides.pop("refine", None)
|
||||
if isinstance(refine, Mapping):
|
||||
for key in _LTX2_REFINE_FLAT_KEYS:
|
||||
for key in REFINE_FLAT_KEYS:
|
||||
if key in refine:
|
||||
kwargs[f"ltx2_refine_{key}"] = refine[key]
|
||||
kwargs.update(preset_overrides)
|
||||
@@ -326,13 +311,10 @@ def request_to_sampling_param(
|
||||
) -> SamplingParam:
|
||||
if request.plan is not None:
|
||||
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
|
||||
if request.state is not None:
|
||||
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
if request.state is not None:
|
||||
_validate_continuation_state(request.state)
|
||||
sampling_param.continuation_state = request.state
|
||||
if request.output.return_state:
|
||||
sampling_param.return_continuation_state = True
|
||||
updates = explicit_request_updates(request)
|
||||
|
||||
for key, value in updates.items():
|
||||
@@ -375,14 +357,9 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
|
||||
|
||||
|
||||
def _compile_config_to_torch_kwargs(compile_config: CompileConfig, ) -> dict[str, Any]:
|
||||
"""Flatten typed ``CompileConfig`` back to a ``torch_compile_kwargs``
|
||||
dict that the legacy ``FastVideoArgs`` path still expects.
|
||||
|
||||
Typed first-class fields (:attr:`backend`, :attr:`fullgraph`,
|
||||
:attr:`mode`, :attr:`dynamic`) are only emitted when the user set
|
||||
them explicitly (non-``None``). ``extras`` is merged on top for any
|
||||
uncommon kwargs.
|
||||
"""
|
||||
"""Flatten typed ``CompileConfig`` back to the legacy
|
||||
``torch_compile_kwargs`` dict, emitting only explicitly-set typed
|
||||
fields and merging ``extras`` on top."""
|
||||
out: dict[str, Any] = {}
|
||||
for key in _COMPILE_TYPED_KEYS:
|
||||
value = getattr(compile_config, key)
|
||||
@@ -553,37 +530,6 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
|
||||
|
||||
_SCHEMA_DEFAULT_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
|
||||
|
||||
_KNOWN_CONTINUATION_KINDS: set[str] = set()
|
||||
|
||||
|
||||
def register_continuation_kind(kind: str) -> None:
|
||||
"""Register a :class:`ContinuationState.kind` as recognized.
|
||||
|
||||
PR 7 wires the envelope through; per-kind payload deserializers live
|
||||
with each model family (e.g. ``fastvideo.pipelines.basic.ltx2.
|
||||
continuation.LTX2ContinuationState``). The registry lets the
|
||||
public-API compat layer validate the kind early, before the state
|
||||
reaches the pipeline.
|
||||
"""
|
||||
if not isinstance(kind, str) or not kind:
|
||||
raise ValueError("ContinuationState kind must be a non-empty string")
|
||||
_KNOWN_CONTINUATION_KINDS.add(kind)
|
||||
|
||||
|
||||
def _validate_continuation_state(state: ContinuationState) -> None:
|
||||
if not isinstance(state.kind, str) or not state.kind:
|
||||
raise ValueError("GenerationRequest.state.kind must be a non-empty string; got "
|
||||
f"{state.kind!r}")
|
||||
if not isinstance(state.payload, Mapping):
|
||||
raise ValueError(f"GenerationRequest.state.payload must be a mapping; got "
|
||||
f"{type(state.payload).__name__}")
|
||||
if state.kind not in _KNOWN_CONTINUATION_KINDS:
|
||||
known = sorted(_KNOWN_CONTINUATION_KINDS)
|
||||
raise ValueError(f"Unknown ContinuationState kind {state.kind!r}; registered "
|
||||
f"kinds: {known}. Import the model family that owns this kind "
|
||||
"(e.g. `import fastvideo.pipelines.basic.ltx2.continuation`) "
|
||||
"to register it, or drop the state field.")
|
||||
|
||||
|
||||
def _fan_out_batched_input_value(
|
||||
source_request: GenerationRequest,
|
||||
@@ -617,7 +563,6 @@ __all__ = [
|
||||
"load_generator_config_from_file",
|
||||
"normalize_generation_request",
|
||||
"normalize_generator_config",
|
||||
"register_continuation_kind",
|
||||
"request_to_pipeline_overrides",
|
||||
"request_to_sampling_param",
|
||||
]
|
||||
|
||||
@@ -15,7 +15,6 @@ class GenerationResult:
|
||||
samples: Any | None = None
|
||||
frames: Any | None = None
|
||||
audio: Any | None = None
|
||||
audio_sample_rate: int | None = None
|
||||
size: tuple[int, int, int] | None = None
|
||||
generation_time: float | None = None
|
||||
logging_info: Any | None = None
|
||||
@@ -45,7 +44,6 @@ class GenerationResult:
|
||||
"samples",
|
||||
"frames",
|
||||
"audio",
|
||||
"audio_sample_rate",
|
||||
"size",
|
||||
"generation_time",
|
||||
"logging_info",
|
||||
@@ -64,7 +62,6 @@ class GenerationResult:
|
||||
samples=result.get("samples"),
|
||||
frames=result.get("frames"),
|
||||
audio=result.get("audio"),
|
||||
audio_sample_rate=result.get("audio_sample_rate"),
|
||||
size=result.get("size"),
|
||||
generation_time=result.get("generation_time"),
|
||||
logging_info=result.get("logging_info"),
|
||||
@@ -83,7 +80,6 @@ class GenerationResult:
|
||||
"samples": self.samples,
|
||||
"frames": self.frames,
|
||||
"audio": self.audio,
|
||||
"audio_sample_rate": self.audio_sample_rate,
|
||||
"size": self.size,
|
||||
"generation_time": self.generation_time,
|
||||
"logging_info": self.logging_info,
|
||||
|
||||
@@ -1,16 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import StoreBoolean
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -97,13 +92,9 @@ class SamplingParam:
|
||||
movement_distance: float | None = None
|
||||
camera_rotation: str | None = None
|
||||
|
||||
# LTX-2 multi-modal CFG and STG.
|
||||
# cfg_scale defaults are 1.0 (CFG off) so ``ForwardBatch.__post_init__``
|
||||
# doesn't force ``do_classifier_free_guidance`` on non-LTX-2 models that
|
||||
# never override these fields. LTX-2 presets that need text-CFG on set
|
||||
# them in their ``defaults`` dict (e.g. ``ltx2_base``).
|
||||
ltx2_cfg_scale_video: float = 1.0
|
||||
ltx2_cfg_scale_audio: float = 1.0
|
||||
# LTX2 multi-modal CFG and STG
|
||||
ltx2_cfg_scale_video: float = 3.0
|
||||
ltx2_cfg_scale_audio: float = 7.0
|
||||
ltx2_modality_scale_video: float = 3.0
|
||||
ltx2_modality_scale_audio: float = 3.0
|
||||
ltx2_rescale_scale: float = 0.7
|
||||
@@ -112,39 +103,6 @@ class SamplingParam:
|
||||
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
|
||||
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
|
||||
|
||||
# Stable Audio (T2A): clip start/end in seconds. Honored by
|
||||
# `StableAudioConditioningStage` + `StableAudioDecodingStage`. Other
|
||||
# families ignore them.
|
||||
audio_start_in_s: float | None = None
|
||||
audio_end_in_s: float | None = None
|
||||
|
||||
# Stable Audio audio-to-audio (variation):
|
||||
# `init_audio` -- a path or `[B, C, samples]` waveform at the model
|
||||
# sample rate; the pipeline encodes it via the VAE
|
||||
# and uses it as the starting latent.
|
||||
# `init_audio_strength` -- 0..1, higher = closer to the reference
|
||||
# (matches the convention of Stability's
|
||||
# commercial Stable Audio 2.0 UI). 1.0 ~=
|
||||
# VAE round-trip, 0.0 ~= plain T2A.
|
||||
# `init_noise_level` -- legacy raw `sigma_max` override (0.3..500,
|
||||
# higher = more freedom). Kept for callers
|
||||
# that already use it; prefer `init_audio_strength`.
|
||||
init_audio: Any = None
|
||||
init_audio_strength: float | None = None
|
||||
init_noise_level: float | None = None
|
||||
|
||||
# Stable Audio inpainting (RePaint-style): `inpaint_audio` is the
|
||||
# reference clip, `inpaint_mask` is a [samples] tensor in {0, 1} where
|
||||
# 1 means *keep the reference* and 0 means *regenerate*.
|
||||
inpaint_audio: Any = None
|
||||
inpaint_mask: Any = None
|
||||
|
||||
# Continuation state carried across streaming/multi-segment calls.
|
||||
continuation_state: ContinuationState | None = None
|
||||
# When True, the pipeline returns a ContinuationState on the result so
|
||||
# the caller can resume from the generated segment.
|
||||
return_continuation_state: bool = False
|
||||
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
@@ -169,7 +127,7 @@ class SamplingParam:
|
||||
self.__post_init__()
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> SamplingParam:
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
sampling_param = cls._from_preset(model_path)
|
||||
if sampling_param is not None:
|
||||
return sampling_param
|
||||
@@ -185,7 +143,7 @@ class SamplingParam:
|
||||
def _from_preset(
|
||||
cls,
|
||||
model_path: str,
|
||||
) -> SamplingParam | None:
|
||||
) -> "SamplingParam | None":
|
||||
"""Build a SamplingParam from preset defaults.
|
||||
|
||||
Returns ``None`` when no preset is configured for
|
||||
|
||||
@@ -41,12 +41,6 @@ class CompileConfig:
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
text_encoder_enabled: bool | None = None
|
||||
"""Whether ``torch.compile`` is applied to the text encoder. ``None``
|
||||
keeps the runtime default. The public ``FastVideoArgs`` adapter does
|
||||
not yet consume this flag; reserved so the realtime runtime upstream
|
||||
(PR 7.6) has a typed home for its ``enable_torch_compile_text_encoder``
|
||||
kwarg without routing through ``pipeline.experimental``."""
|
||||
backend: str | None = None
|
||||
fullgraph: bool | None = None
|
||||
mode: str | None = None
|
||||
|
||||
@@ -48,6 +48,8 @@ class ModelConfig:
|
||||
for key, value in source_model_dict.items():
|
||||
if key in valid_fields:
|
||||
setattr(arch_config, key, value)
|
||||
else:
|
||||
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
|
||||
|
||||
if hasattr(arch_config, "__post_init__"):
|
||||
arch_config.__post_init__()
|
||||
|
||||
@@ -5,13 +5,11 @@ from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"StableAudioConfig"
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig"
|
||||
]
|
||||
|
||||
@@ -50,8 +50,7 @@ class CosmosArchConfig(DiTArchConfig):
|
||||
})
|
||||
|
||||
# Cosmos-specific config parameters based on transformer_cosmos.py
|
||||
# in_channels includes the condition_mask channel (16 latent + 1 cond = 17)
|
||||
in_channels: int = 17
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128
|
||||
|
||||
@@ -1,76 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the Stable Audio Open 1.0 DiT.
|
||||
|
||||
Note: the SA pipeline bypasses the standard `ComposedPipelineBase`
|
||||
component loader because the published HF repo ships a single monolithic
|
||||
`model.safetensors` (no Diffusers-style `model_index.json` or
|
||||
per-subfolder layout). The arch fields and `param_names_mapping` here
|
||||
document the architecture and key remap so the same conventions used by
|
||||
the rest of the DiT family apply (FSDP shard conditions, supported
|
||||
attention backends, future loader integrations) — they are not currently
|
||||
consumed by `fastvideo/models/loader/fsdp_load.py` for SA.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
# Matches `transformer.layers.{i}` in the SA DiT module tree.
|
||||
parts = n.split(".")
|
||||
return (len(parts) >= 3 and parts[-3] == "transformer" and parts[-2] == "layers" and parts[-1].isdigit())
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer])
|
||||
|
||||
# SA's checkpoint is `stable_audio_tools` raw format (not Diffusers),
|
||||
# so the only remaps are: strip the `model.model.` host-pipeline
|
||||
# prefix, and rename `nn.LayerNorm`'s `gamma`/`beta` to torch's
|
||||
# canonical `weight`/`bias`. Linear / cross-attention naming already
|
||||
# matches FastVideo's conventions, so no further remap is needed.
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^model\.model\.(.*?)\.gamma$": r"\1.weight",
|
||||
r"^model\.model\.(.*?)\.beta$": r"\1.bias",
|
||||
r"^model\.model\.(.*)$": r"\1",
|
||||
})
|
||||
|
||||
# SA only supports backends compatible with single-GPU LocalAttention.
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
# Architecture constants (from the published `model_config.json` for
|
||||
# `stabilityai/stable-audio-open-1.0`).
|
||||
io_channels: int = 64
|
||||
embed_dim: int = 1536
|
||||
depth: int = 24
|
||||
num_attention_heads: int = 24
|
||||
cond_token_dim: int = 768
|
||||
global_cond_dim: int = 1536
|
||||
project_cond_tokens: bool = False
|
||||
project_global_cond: bool = True
|
||||
# Set to "ln" to wrap attention Q/K in LayerNorm (used by
|
||||
# `stable-audio-open-small`; absent in the 1.0 base).
|
||||
qk_norm: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.hidden_size = self.embed_dim
|
||||
self.in_channels = self.io_channels
|
||||
self.out_channels = self.io_channels
|
||||
self.num_channels_latents = self.io_channels
|
||||
self.attention_head_dim = self.embed_dim // self.num_attention_heads
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=StableAudioArchConfig)
|
||||
|
||||
prefix: str = "StableAudio"
|
||||
@@ -7,12 +7,9 @@ from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig)
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig"
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig"
|
||||
]
|
||||
|
||||
@@ -1,80 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the Stable Audio Open 1.0 multi-conditioner.
|
||||
|
||||
The conditioner bundles three sub-conditioners — a T5 text encoder
|
||||
(prompt) and two NumberConditioners (`seconds_start` / `seconds_total`)
|
||||
— into the (cross_attn_cond, cross_attn_mask, global_embed) triple the
|
||||
DiT consumes. The architecture is fully specified by the official
|
||||
`stable_audio_tools` `MultiConditioner` config; the constants here
|
||||
mirror that.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig
|
||||
from fastvideo.configs.models.encoders.base import (EncoderArchConfig, EncoderConfig)
|
||||
|
||||
|
||||
def _default_configs() -> list[dict]:
|
||||
"""Default = `stable-audio-open-1.0`'s three sub-conditioners."""
|
||||
return [
|
||||
{
|
||||
"id": "prompt",
|
||||
"type": "t5",
|
||||
"config": {
|
||||
"t5_model_name": "t5-base",
|
||||
"max_length": 128
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "seconds_start",
|
||||
"type": "number",
|
||||
"config": {
|
||||
"min_val": 0,
|
||||
"max_val": 512
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "seconds_total",
|
||||
"type": "number",
|
||||
"config": {
|
||||
"min_val": 0,
|
||||
"max_val": 512
|
||||
}
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioConditionerArchConfig(EncoderArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["StableAudioMultiConditioner"])
|
||||
|
||||
# Shared embedding width across all sub-conditioners (T5 last-hidden
|
||||
# dim and NumberEmbedder feature dim both = `cond_dim`).
|
||||
cond_dim: int = 768
|
||||
|
||||
# Sub-conditioner identifiers. Order in `cross_attention_cond_ids`
|
||||
# is the concat order for the cross-attn token sequence; order in
|
||||
# `global_cond_ids` is the concat order for the global FiLM-style
|
||||
# embedding.
|
||||
cross_attention_cond_ids: tuple[str, ...] = ("prompt", "seconds_start", "seconds_total")
|
||||
global_cond_ids: tuple[str, ...] = ("seconds_start", "seconds_total")
|
||||
|
||||
# Per-sub-conditioner spec list (mirrors upstream
|
||||
# `model_config.json.model.conditioning.configs`). Each entry is
|
||||
# `{"id": ..., "type": "t5"|"number", "config": {...}}`. The default
|
||||
# matches `stable-audio-open-1.0`; SA-small overrides via the
|
||||
# `conditioner/config.json` shipped in the converted repo.
|
||||
configs: list = field(default_factory=_default_configs)
|
||||
|
||||
# Match official `stable_audio_tools/models/conditioners.py:334`:
|
||||
# T5 is loaded directly in fp16.
|
||||
t5_dtype: str = "float16"
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioConditionerConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=StableAudioConditionerArchConfig)
|
||||
|
||||
prefix: str = "stable_audio_conditioner"
|
||||
@@ -41,14 +41,6 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
text_len: int = 512
|
||||
dtype: str | None = None
|
||||
gradient_checkpointing: bool = False
|
||||
# Extra fields present in upstream HF T5Config but unused by FastVideo's
|
||||
# encoder. Declared here so `update_model_arch` doesn't reject them when
|
||||
# loading repos like `stabilityai/stable-audio-open-1.0` that ship the
|
||||
# full HF config.
|
||||
n_positions: int = 512
|
||||
decoder_start_token_id: int = 0
|
||||
output_past: bool = True
|
||||
task_specific_params: dict | None = None
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
|
||||
@@ -5,7 +5,6 @@ from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
from fastvideo.configs.models.vaes.oobleck import OobleckVAEArchConfig, OobleckVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
__all__ = [
|
||||
@@ -17,6 +16,4 @@ __all__ = [
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
]
|
||||
|
||||
@@ -1,68 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the Stable Audio Open 1.0 "Oobleck" VAE.
|
||||
|
||||
Mirrors the per-channel `vae/config.json` shipped in
|
||||
`stabilityai/stable-audio-open-1.0` 1:1 (see
|
||||
`fastvideo/models/vaes/oobleck.py::OobleckVAE.from_pretrained`, which
|
||||
constructs the VAE from these fields). Inherits the FastVideo VAEConfig
|
||||
base so the standard `load_encoder` / `load_decoder` flags + tiling
|
||||
knobs apply.
|
||||
|
||||
Naming: the VAE architecture is officially "Oobleck" (per Stability
|
||||
AI's stable-audio-tools) — the surrounding model family is "Stable
|
||||
Audio Open 1.0". This config is named after the architecture
|
||||
(`OobleckVAEConfig`) since the same VAE is shared across Stable Audio
|
||||
checkpoints; downstream pipelines reference it by its arch name, not
|
||||
by a host-pipeline name.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class OobleckVAEArchConfig(VAEArchConfig):
|
||||
"""Stable Audio Open 1.0 VAE architecture constants."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["AutoencoderOobleck"])
|
||||
|
||||
# From stabilityai/stable-audio-open-1.0/vae/config.json.
|
||||
encoder_hidden_size: int = 128
|
||||
downsampling_ratios: list[int] = field(default_factory=lambda: [2, 4, 4, 8, 8])
|
||||
channel_multiples: list[int] = field(default_factory=lambda: [1, 2, 4, 8, 16])
|
||||
decoder_channels: int = 128
|
||||
decoder_input_channels: int = 64
|
||||
audio_channels: int = 2 # stereo
|
||||
sampling_rate: int = 44100
|
||||
|
||||
|
||||
@dataclass
|
||||
class OobleckVAEConfig(VAEConfig):
|
||||
"""FastVideo VAE config wrapping the Oobleck arch.
|
||||
|
||||
Audio VAEs don't use the temporal/spatial tiling defaults that the
|
||||
base VAEConfig is shaped for (those exist for video VAEs); they are
|
||||
retained but irrelevant for audio.
|
||||
"""
|
||||
|
||||
arch_config: VAEArchConfig = field(default_factory=OobleckVAEArchConfig)
|
||||
|
||||
# Audio is 1-D, so the video-VAE tiling defaults are inert. Disable
|
||||
# them so callers don't accidentally trip on tile-stride math built
|
||||
# for spatial tensors.
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
# Where the FastVideo loader / pipeline-glue wrapper should fetch
|
||||
# weights from when no local path is supplied. Gated repo — caller's
|
||||
# HF token must have accepted terms on
|
||||
# https://huggingface.co/stabilityai/stable-audio-open-1.0.
|
||||
pretrained_path: str = "stabilityai/stable-audio-open-1.0"
|
||||
pretrained_subfolder: str = "vae"
|
||||
# Match official `stable_audio_tools`: VAE runs in fp16 (the
|
||||
# `pretransform.model_half` path in
|
||||
# `stable_audio_tools/models/pretransforms.py`).
|
||||
pretrained_dtype: str = "float16"
|
||||
@@ -11,34 +11,10 @@ import torch
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5ArchConfig(T5ArchConfig):
|
||||
"""T5 arch that pads tokenizer output to ``max_length``.
|
||||
|
||||
LongCat's denoising stage concatenates positive and negative
|
||||
attention masks along the batch dimension for CFG, which requires
|
||||
uniform seq length. The shared :class:`T5ArchConfig` dropped the
|
||||
``"padding": "max_length"`` tokenizer kwarg so other DiTs could run
|
||||
with variable-length masks; LongCat still needs the uniform
|
||||
contract.
|
||||
"""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT5Config(T5Config):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LongCatT5ArchConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields."""
|
||||
@@ -127,9 +103,8 @@ class LongCatT2V480PConfig(PipelineConfig):
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# UMT5 uses T5-like config; postprocess pads to 512. LongCatT5Config
|
||||
# restores ``padding="max_length"`` for the CFG concat contract.
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (LongCatT5Config(), ))
|
||||
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (longcat_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (umt5_postprocess_text, ))
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""`PipelineConfig` for Stable Audio Open 1.0."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import StableAudioConfig
|
||||
from fastvideo.configs.models.vaes import OobleckVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioT2AConfig(PipelineConfig):
|
||||
"""Stable Audio Open 1.0 pipeline config."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=StableAudioConfig)
|
||||
# Standard `TransformerLoader` reads `dit_precision`; default in
|
||||
# `PipelineConfig` is bf16, but we want fp16 to match official.
|
||||
dit_precision: str = "fp16"
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=OobleckVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# `StableAudioMultiConditioner` owns its own T5; zero out the
|
||||
# parent's text-encoder slots so the length-equality validator passes.
|
||||
text_encoder_configs: tuple = field(default_factory=tuple)
|
||||
preprocess_text_funcs: tuple = field(default_factory=tuple)
|
||||
postprocess_text_funcs: tuple = field(default_factory=tuple)
|
||||
|
||||
num_inference_steps: int = 100
|
||||
guidance_scale: float = 7.0
|
||||
audio_end_in_s: float = 10.0 # short-clip default
|
||||
audio_start_in_s: float = 0.0
|
||||
sampling_rate: int = 44100
|
||||
audio_channels: int = 2
|
||||
# Stable Audio Open 1.0 was trained at a fixed 2,097,152-sample
|
||||
# window (= 2097152 / 44100 ≈ 47.55s). Anything past this is
|
||||
# silently truncated by the post-decode slice — validate up-front.
|
||||
sample_size: int = 2097152
|
||||
max_audio_duration_s: float = 2097152 / 44100
|
||||
|
||||
# Match the official `stable_audio_tools` defaults (`model_half=True`
|
||||
# in `run_gradio.py`), which loads the DiT, VAE, and T5 in fp16 and
|
||||
# wraps T5 forward in `autocast(fp16)`. fp16 is also a hard
|
||||
# requirement for FlashAttention-2 / FA-3.
|
||||
precision: str = "fp16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# A2A needs encode; load both halves for either path.
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioOpenSmallConfig(StableAudioT2AConfig):
|
||||
"""`stable-audio-open-small` overrides: shorter training window
|
||||
(524288 samples ≈ 11.89s @ 44.1 kHz) and a faster default sampler
|
||||
config carried by the small preset.
|
||||
"""
|
||||
|
||||
sample_size: int = 524288
|
||||
max_audio_duration_s: float = 524288 / 44100
|
||||
audio_end_in_s: float = 6.0 # short-clip default suitable for the small window
|
||||
@@ -1,31 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.entrypoints.streaming.server import build_app, run_server
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
BlobStore,
|
||||
InMemoryBlobStore,
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.server import run_server
|
||||
|
||||
__all__ = [
|
||||
"BlobStore",
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionState",
|
||||
"SessionStore",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
__all__ = ["run_server"]
|
||||
|
||||
@@ -1,252 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""JSON WebSocket protocol schemas for the streaming server.
|
||||
|
||||
Every control message shares the envelope ``{"type": <str>, ...}``.
|
||||
Pydantic models live here so the server can parse / validate incoming
|
||||
frames and emit well-typed outgoing frames without hand-rolled dicts.
|
||||
|
||||
The message catalogue matches the contract in
|
||||
``docs/design/server_contracts/streaming.md``; additions must land in
|
||||
both places in the same PR.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated, Any, Literal, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client → server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SessionInitV2(BaseModel):
|
||||
"""Opening frame the client sends after the WebSocket handshake."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
type: Literal["session_init_v2"]
|
||||
client_id: str | None = None
|
||||
preset: str | None = None
|
||||
preset_label: str | None = None
|
||||
curated_prompts: list[str] = Field(default_factory=list)
|
||||
initial_image: dict[str, Any] | None = None
|
||||
enhancement_enabled: bool = False
|
||||
auto_extension_enabled: bool = False
|
||||
loop_generation_enabled: bool = False
|
||||
single_clip_mode: bool = False
|
||||
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
|
||||
continuation_state: dict[str, Any] | None = None
|
||||
"""Optional ``{kind, payload}`` dict; hydrated into
|
||||
:class:`fastvideo.api.ContinuationState` server-side."""
|
||||
|
||||
|
||||
class SegmentPromptSource(BaseModel):
|
||||
"""Request a new segment using a specific prompt."""
|
||||
|
||||
type: Literal["segment_prompt_source"]
|
||||
prompt: str
|
||||
negative_prompt: str | None = None
|
||||
source: Literal["curated", "enhanced", "user", "auto_extension"] = "user"
|
||||
seed: int | None = None
|
||||
num_inference_steps: int | None = None
|
||||
guidance_scale: float | None = None
|
||||
|
||||
|
||||
class SeedPromptsUpdated(BaseModel):
|
||||
type: Literal["seed_prompts_updated"]
|
||||
seed_prompts: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class EnhancementUpdated(BaseModel):
|
||||
type: Literal["enhancement_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutoExtensionUpdated(BaseModel):
|
||||
type: Literal["auto_extension_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class LoopGenerationUpdated(BaseModel):
|
||||
type: Literal["loop_generation_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class GenerationPausedUpdated(BaseModel):
|
||||
type: Literal["generation_paused_updated"]
|
||||
paused: bool
|
||||
|
||||
|
||||
class SnapshotState(BaseModel):
|
||||
"""Request the current ``ContinuationState`` for export."""
|
||||
|
||||
type: Literal["snapshot_state"]
|
||||
|
||||
|
||||
ClientMessage = Annotated[
|
||||
Union[ # noqa: UP007 - Annotated requires Union for discriminator
|
||||
SessionInitV2,
|
||||
SegmentPromptSource,
|
||||
SeedPromptsUpdated,
|
||||
EnhancementUpdated,
|
||||
AutoExtensionUpdated,
|
||||
LoopGenerationUpdated,
|
||||
GenerationPausedUpdated,
|
||||
SnapshotState,
|
||||
],
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Server → client
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class QueueStatus(BaseModel):
|
||||
type: Literal["queue_status"] = "queue_status"
|
||||
position: int
|
||||
queue_depth: int
|
||||
|
||||
|
||||
class GpuAssigned(BaseModel):
|
||||
type: Literal["gpu_assigned"] = "gpu_assigned"
|
||||
gpu_id: int
|
||||
session_timeout: int
|
||||
|
||||
|
||||
class Ltx2StreamStart(BaseModel):
|
||||
type: Literal["ltx2_stream_start"] = "ltx2_stream_start"
|
||||
preset: str | None = None
|
||||
width: int
|
||||
height: int
|
||||
fps: int
|
||||
num_frames: int
|
||||
|
||||
|
||||
class Ltx2SegmentStart(BaseModel):
|
||||
type: Literal["ltx2_segment_start"] = "ltx2_segment_start"
|
||||
segment_idx: int
|
||||
prompt: str
|
||||
total_steps: int
|
||||
|
||||
|
||||
class StepComplete(BaseModel):
|
||||
type: Literal["step_complete"] = "step_complete"
|
||||
segment_idx: int
|
||||
step: int
|
||||
total_steps: int
|
||||
stage: str = "denoise"
|
||||
|
||||
|
||||
class MediaInit(BaseModel):
|
||||
"""Descriptor for the fMP4 initialization segment that follows."""
|
||||
|
||||
type: Literal["media_init"] = "media_init"
|
||||
segment_idx: int
|
||||
mime: str = "video/mp4; codecs=\"avc1.64001f, mp4a.40.2\""
|
||||
stream_id: str
|
||||
mode: Literal["av_fmp4"] = "av_fmp4"
|
||||
|
||||
|
||||
class MediaSegmentComplete(BaseModel):
|
||||
type: Literal["media_segment_complete"] = "media_segment_complete"
|
||||
segment_idx: int
|
||||
stream_id: str
|
||||
chunks: int
|
||||
duration_ms: float | None = None
|
||||
pts_base_ms: float | None = None
|
||||
|
||||
|
||||
class Ltx2SegmentComplete(BaseModel):
|
||||
type: Literal["ltx2_segment_complete"] = "ltx2_segment_complete"
|
||||
segment_idx: int
|
||||
generation_time_ms: float
|
||||
e2e_latency_ms: float | None = None
|
||||
|
||||
|
||||
class Ltx2StreamComplete(BaseModel):
|
||||
type: Literal["ltx2_stream_complete"] = "ltx2_stream_complete"
|
||||
reason: Literal["segment_cap", "stop_requested", "error"] = "stop_requested"
|
||||
|
||||
|
||||
class SessionTimeout(BaseModel):
|
||||
type: Literal["session_timeout"] = "session_timeout"
|
||||
timeout_seconds: int
|
||||
|
||||
|
||||
class ContinuationStateSnapshot(BaseModel):
|
||||
type: Literal["continuation_state_snapshot"] = "continuation_state_snapshot"
|
||||
state: dict[str, Any]
|
||||
"""``{kind, payload}`` dict matching
|
||||
:class:`fastvideo.api.ContinuationState`."""
|
||||
|
||||
|
||||
class ErrorMessage(BaseModel):
|
||||
type: Literal["error"] = "error"
|
||||
code: Literal[
|
||||
"session_rejected",
|
||||
"invalid_message",
|
||||
"preset_mismatch",
|
||||
"gpu_unavailable",
|
||||
"worker_failed",
|
||||
"upstream_timeout",
|
||||
"internal_error",
|
||||
] = "internal_error"
|
||||
message: str
|
||||
retryable: bool = False
|
||||
|
||||
|
||||
ServerMessage = Union[ # noqa: UP007 - pydantic Union handling
|
||||
QueueStatus,
|
||||
GpuAssigned,
|
||||
Ltx2StreamStart,
|
||||
Ltx2SegmentStart,
|
||||
StepComplete,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2StreamComplete,
|
||||
SessionTimeout,
|
||||
ContinuationStateSnapshot,
|
||||
ErrorMessage,
|
||||
]
|
||||
|
||||
|
||||
def parse_client_message(raw: dict[str, Any]) -> ClientMessage:
|
||||
"""Parse an incoming WebSocket dict into a typed client message.
|
||||
|
||||
Unknown ``type`` values raise :class:`pydantic.ValidationError`; the
|
||||
server handler turns that into an ``error`` frame with
|
||||
``code="invalid_message"``.
|
||||
"""
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
return TypeAdapter(ClientMessage).validate_python(raw)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AutoExtensionUpdated",
|
||||
"ClientMessage",
|
||||
"ContinuationStateSnapshot",
|
||||
"EnhancementUpdated",
|
||||
"ErrorMessage",
|
||||
"GenerationPausedUpdated",
|
||||
"GpuAssigned",
|
||||
"Ltx2SegmentComplete",
|
||||
"Ltx2SegmentStart",
|
||||
"Ltx2StreamComplete",
|
||||
"Ltx2StreamStart",
|
||||
"LoopGenerationUpdated",
|
||||
"MediaInit",
|
||||
"MediaSegmentComplete",
|
||||
"QueueStatus",
|
||||
"SeedPromptsUpdated",
|
||||
"SegmentPromptSource",
|
||||
"ServerMessage",
|
||||
"SessionInitV2",
|
||||
"SessionTimeout",
|
||||
"SnapshotState",
|
||||
"StepComplete",
|
||||
"parse_client_message",
|
||||
]
|
||||
@@ -1,531 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-generator FastAPI + WebSocket streaming server."""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
AutoExtensionUpdated,
|
||||
ContinuationStateSnapshot,
|
||||
EnhancementUpdated,
|
||||
ErrorMessage,
|
||||
GenerationPausedUpdated,
|
||||
GpuAssigned,
|
||||
LoopGenerationUpdated,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamComplete,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SeedPromptsUpdated,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image, )
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
|
||||
from fastvideo.api.schema import ServeConfig
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# RFC 6455 WebSocket close codes used by the server.
|
||||
_WS_CLOSE_UNSUPPORTED_DATA = 1003
|
||||
_WS_CLOSE_TRY_AGAIN_LATER = 1013
|
||||
|
||||
|
||||
class _GeneratorProto(Protocol):
|
||||
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
|
||||
|
||||
def generate(self, request: GenerationRequest) -> Any:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerState:
|
||||
serve_config: ServeConfig
|
||||
generator: _GeneratorProto
|
||||
sessions: SessionManager
|
||||
session_store: SessionStore
|
||||
|
||||
|
||||
def build_app(
|
||||
serve_config: ServeConfig,
|
||||
generator: _GeneratorProto,
|
||||
*,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI app used by :func:`run_server`.
|
||||
|
||||
Exposed so tests can drive the WebSocket endpoint in-process via
|
||||
``starlette.testclient.TestClient(app).websocket_connect(...)``.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
|
||||
"server; got None. Add a `streaming:` block to your serve config.")
|
||||
|
||||
sessions = SessionManager(
|
||||
segment_cap=serve_config.streaming.generation_segment_cap,
|
||||
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
|
||||
)
|
||||
state = ServerState(
|
||||
serve_config=serve_config,
|
||||
generator=generator,
|
||||
sessions=sessions,
|
||||
session_store=session_store or InMemorySessionStore(),
|
||||
)
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming")
|
||||
|
||||
@app.get("/health")
|
||||
async def _health() -> JSONResponse:
|
||||
return JSONResponse({
|
||||
"status": "ok",
|
||||
"sessions": len(state.sessions),
|
||||
"stream_mode": state.serve_config.streaming.stream_mode,
|
||||
})
|
||||
|
||||
@app.websocket("/v1/stream")
|
||||
async def _stream(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
try:
|
||||
session = state.sessions.create()
|
||||
except SessionRejected as exc:
|
||||
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
|
||||
return
|
||||
|
||||
try:
|
||||
await _handle_session(websocket, session, state)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("session %s: client disconnected", session.id[:8])
|
||||
except Exception: # pragma: no cover - defensive catch-all
|
||||
logger.exception("session %s: unhandled error", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
finally:
|
||||
_cleanup_session(session, state)
|
||||
|
||||
app.state.server_state = state
|
||||
return app
|
||||
|
||||
|
||||
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
|
||||
"""Launch the streaming server.
|
||||
|
||||
Boots a :class:`fastvideo.VideoGenerator` from
|
||||
``serve_config.generator`` unless ``generator`` is provided, then
|
||||
serves ``build_app(...)`` via uvicorn.
|
||||
"""
|
||||
def run_server(serve_config: ServeConfig) -> None:
|
||||
"""Launch the streaming (WebSocket / Dynamo) server."""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
|
||||
"got None. Add a `streaming:` block to your serve config.")
|
||||
|
||||
import uvicorn
|
||||
|
||||
if generator is None:
|
||||
from fastvideo import VideoGenerator # lazy to avoid boot cost
|
||||
|
||||
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
|
||||
app = build_app(serve_config, generator)
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
)
|
||||
|
||||
|
||||
async def _handle_session(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
init = await _read_init_message(websocket, session, state)
|
||||
if init is None:
|
||||
return
|
||||
|
||||
await _apply_session_init(session, init, state)
|
||||
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
|
||||
session.transition(SessionState.GPU_BINDING)
|
||||
await _send_json(websocket, GpuAssigned(
|
||||
gpu_id=0,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
session.transition(SessionState.ACTIVE)
|
||||
await _send_json(websocket, _build_stream_start(session, state))
|
||||
|
||||
try:
|
||||
await _run_segment_loop(websocket, session, state)
|
||||
finally:
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
|
||||
|
||||
|
||||
async def _read_init_message(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> SessionInitV2 | None:
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: init timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return None
|
||||
except WebSocketDisconnect:
|
||||
return None
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
|
||||
return None
|
||||
if not isinstance(parsed, SessionInitV2):
|
||||
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
|
||||
return None
|
||||
return parsed
|
||||
|
||||
|
||||
async def _reject_init(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
message: str,
|
||||
close_reason: str,
|
||||
) -> None:
|
||||
await _send_error(websocket, "invalid_message", message, retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.REJECTED)
|
||||
|
||||
|
||||
async def _apply_session_init(
|
||||
session: Session,
|
||||
init: SessionInitV2,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
session.client_id = init.client_id
|
||||
session.preset = init.preset
|
||||
session.preset_label = init.preset_label
|
||||
session.curated_prompts = list(init.curated_prompts)
|
||||
session.enhancement_enabled = init.enhancement_enabled
|
||||
session.auto_extension_enabled = init.auto_extension_enabled
|
||||
session.loop_generation_enabled = init.loop_generation_enabled
|
||||
session.single_clip_mode = init.single_clip_mode
|
||||
session.stream_mode = init.stream_mode
|
||||
|
||||
if init.initial_image is not None:
|
||||
# Decode + disk write off the event loop; payload is up to 32 MiB.
|
||||
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
|
||||
if image is not None:
|
||||
session.metadata["session_init_image"] = image.path
|
||||
|
||||
if init.continuation_state is not None:
|
||||
session.continuation_state = _coerce_state(init.continuation_state)
|
||||
if session.continuation_state is not None:
|
||||
state.session_store.store(session.id, session.continuation_state)
|
||||
|
||||
|
||||
async def _run_segment_loop(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
cap = state.sessions.segment_cap
|
||||
while True:
|
||||
if session.segment_cap_reached(cap):
|
||||
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
|
||||
return
|
||||
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: idle timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return
|
||||
except WebSocketDisconnect:
|
||||
return
|
||||
session.touch()
|
||||
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SnapshotState):
|
||||
snap = state.session_store.snapshot(session.id)
|
||||
if snap is None:
|
||||
await _send_error(websocket,
|
||||
"internal_error",
|
||||
"no continuation state available for session",
|
||||
retryable=False)
|
||||
continue
|
||||
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SegmentPromptSource):
|
||||
await _run_segment(websocket, session, state, parsed)
|
||||
continue
|
||||
|
||||
# Silently ignore unknown-but-valid types (additive-evolution
|
||||
# rule in streaming.md).
|
||||
_apply_toggle(session, parsed)
|
||||
|
||||
|
||||
async def _run_segment(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
message: SegmentPromptSource,
|
||||
) -> None:
|
||||
request = _build_generation_request(session, message, state)
|
||||
segment_idx = session.segment_idx
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentStart(
|
||||
segment_idx=segment_idx,
|
||||
prompt=message.prompt,
|
||||
total_steps=request.sampling.num_inference_steps,
|
||||
))
|
||||
|
||||
start = time.perf_counter()
|
||||
loop = asyncio.get_running_loop()
|
||||
# TODO: executor-wrapped generate() cannot be cancelled, so a
|
||||
# client disconnect mid-segment leaves the GPU work running to
|
||||
# completion. Real cancellation needs the generate_async API.
|
||||
try:
|
||||
result = await loop.run_in_executor(None, state.generator.generate, request)
|
||||
except Exception as exc:
|
||||
logger.exception("session %s: generator failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
elapsed_ms = (time.perf_counter() - start) * 1000.0
|
||||
|
||||
frames = _extract_frames(result)
|
||||
if not frames:
|
||||
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
|
||||
# Synchronous generator call has no per-step hook; emit one
|
||||
# terminal StepComplete so observability wiring still sees the
|
||||
# segment finish.
|
||||
total = request.sampling.num_inference_steps
|
||||
await _send_json(websocket, StepComplete(
|
||||
segment_idx=segment_idx,
|
||||
step=total,
|
||||
total_steps=total,
|
||||
stage="denoise",
|
||||
))
|
||||
|
||||
encoder = FragmentedMP4Encoder(
|
||||
width=request.sampling.width,
|
||||
height=request.sampling.height,
|
||||
fps=request.sampling.fps,
|
||||
segment_idx=segment_idx,
|
||||
)
|
||||
chunks_relayed = 0
|
||||
async with encoder:
|
||||
init_sent = False
|
||||
async for chunk in encoder.encode(frames):
|
||||
if chunk.kind == "init":
|
||||
await _send_json(websocket, MediaInit(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=chunk.stream_id,
|
||||
))
|
||||
init_sent = True
|
||||
await websocket.send_bytes(chunk.data)
|
||||
if init_sent and chunk.kind == "media":
|
||||
chunks_relayed += 1
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
MediaSegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=encoder.stream_id,
|
||||
chunks=chunks_relayed,
|
||||
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
|
||||
))
|
||||
|
||||
new_state = _extract_state(result)
|
||||
if new_state is not None:
|
||||
session.continuation_state = new_state
|
||||
state.session_store.store(session.id, new_state)
|
||||
|
||||
session.segment_idx += 1
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ACTIVE)
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
generation_time_ms=elapsed_ms,
|
||||
e2e_latency_ms=elapsed_ms,
|
||||
))
|
||||
|
||||
|
||||
def _build_stream_start(
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> Ltx2StreamStart:
|
||||
default = state.serve_config.default_request
|
||||
return Ltx2StreamStart(
|
||||
preset=session.preset,
|
||||
width=default.sampling.width,
|
||||
height=default.sampling.height,
|
||||
fps=default.sampling.fps,
|
||||
num_frames=default.sampling.num_frames,
|
||||
)
|
||||
|
||||
|
||||
def _build_generation_request(
|
||||
session: Session,
|
||||
message: SegmentPromptSource,
|
||||
state: ServerState,
|
||||
) -> GenerationRequest:
|
||||
# Start from the operator-pinned default_request to pick up the
|
||||
# preset-selected sampling knobs; override with per-message values.
|
||||
base = state.serve_config.default_request
|
||||
sampling_kwargs: dict[str, Any] = {
|
||||
"num_videos_per_prompt":
|
||||
base.sampling.num_videos_per_prompt,
|
||||
"seed":
|
||||
message.seed if message.seed is not None else base.sampling.seed,
|
||||
"num_frames":
|
||||
base.sampling.num_frames,
|
||||
"height":
|
||||
base.sampling.height,
|
||||
"width":
|
||||
base.sampling.width,
|
||||
"fps":
|
||||
base.sampling.fps,
|
||||
"num_inference_steps":
|
||||
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
|
||||
"guidance_scale":
|
||||
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
|
||||
}
|
||||
request = GenerationRequest(
|
||||
prompt=message.prompt,
|
||||
negative_prompt=message.negative_prompt or base.negative_prompt,
|
||||
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
|
||||
sampling=SamplingConfig(**sampling_kwargs),
|
||||
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
|
||||
state=session.continuation_state,
|
||||
)
|
||||
return request
|
||||
|
||||
|
||||
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
|
||||
kind = raw.get("kind")
|
||||
payload = raw.get("payload")
|
||||
if not isinstance(kind, str) or not isinstance(payload, dict):
|
||||
return None
|
||||
return ContinuationState(kind=kind, payload=payload)
|
||||
|
||||
|
||||
def _apply_toggle(session: Session, message: Any) -> None:
|
||||
if isinstance(message, EnhancementUpdated):
|
||||
session.enhancement_enabled = message.enabled
|
||||
elif isinstance(message, AutoExtensionUpdated):
|
||||
session.auto_extension_enabled = message.enabled
|
||||
elif isinstance(message, LoopGenerationUpdated):
|
||||
session.loop_generation_enabled = message.enabled
|
||||
elif isinstance(message, GenerationPausedUpdated):
|
||||
session.generation_paused = message.paused
|
||||
elif isinstance(message, SeedPromptsUpdated):
|
||||
session.curated_prompts = list(message.seed_prompts)
|
||||
|
||||
|
||||
def _extract_frames(result: Any) -> list:
|
||||
if hasattr(result, "frames"):
|
||||
return list(result.frames or [])
|
||||
if isinstance(result, dict):
|
||||
return list(result.get("frames") or [])
|
||||
return []
|
||||
|
||||
|
||||
def _extract_state(result: Any) -> ContinuationState | None:
|
||||
state = getattr(result, "state", None)
|
||||
if state is None and isinstance(result, dict):
|
||||
state = result.get("state")
|
||||
if isinstance(state, ContinuationState):
|
||||
return state
|
||||
if isinstance(state, dict):
|
||||
return _coerce_state(state)
|
||||
return None
|
||||
|
||||
|
||||
async def _send_json(websocket: WebSocket, message: Any) -> None:
|
||||
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
|
||||
await websocket.send_json(payload)
|
||||
|
||||
|
||||
async def _send_error(
|
||||
websocket: WebSocket,
|
||||
code: str,
|
||||
message: str,
|
||||
*,
|
||||
retryable: bool,
|
||||
) -> None:
|
||||
await _send_json(
|
||||
websocket,
|
||||
ErrorMessage(code=code, message=message, retryable=retryable),
|
||||
)
|
||||
|
||||
|
||||
def _cleanup_session(session: Session, state: ServerState) -> None:
|
||||
state.sessions.close(session.id)
|
||||
state.session_store.drop(session.id)
|
||||
init_image_path = session.metadata.get("session_init_image")
|
||||
if isinstance(init_image_path, str):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(init_image_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ServerState",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
raise NotImplementedError("streaming server is not implemented yet")
|
||||
|
||||
@@ -1,214 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-connection session lifecycle for the streaming server.
|
||||
|
||||
Each WebSocket opens exactly one :class:`Session`. :class:`SessionManager`
|
||||
enforces the ``generation_segment_cap`` and ``session_timeout_seconds``
|
||||
budgets from :class:`fastvideo.api.StreamingConfig`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
class SessionState(enum.Enum):
|
||||
"""State-machine positions for a streaming session.
|
||||
|
||||
Transitions are server-owned. See
|
||||
``docs/design/server_contracts/streaming.md`` for the full diagram.
|
||||
"""
|
||||
|
||||
INITIALIZING = "initializing"
|
||||
QUEUED = "queued"
|
||||
GPU_BINDING = "gpu_binding"
|
||||
ACTIVE = "active"
|
||||
COMPLETE = "complete"
|
||||
ERROR = "error"
|
||||
TIMEOUT = "timeout"
|
||||
REJECTED = "rejected"
|
||||
|
||||
|
||||
_VALID_TRANSITIONS: dict[SessionState, frozenset[SessionState]] = {
|
||||
SessionState.INITIALIZING:
|
||||
frozenset({
|
||||
SessionState.QUEUED,
|
||||
SessionState.GPU_BINDING,
|
||||
SessionState.REJECTED,
|
||||
SessionState.ERROR,
|
||||
}),
|
||||
SessionState.QUEUED:
|
||||
frozenset({
|
||||
SessionState.GPU_BINDING,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
SessionState.REJECTED,
|
||||
}),
|
||||
SessionState.GPU_BINDING:
|
||||
frozenset({
|
||||
SessionState.ACTIVE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
}),
|
||||
SessionState.ACTIVE:
|
||||
frozenset({
|
||||
SessionState.ACTIVE,
|
||||
SessionState.COMPLETE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
}),
|
||||
SessionState.COMPLETE:
|
||||
frozenset(),
|
||||
SessionState.ERROR:
|
||||
frozenset(),
|
||||
SessionState.TIMEOUT:
|
||||
frozenset(),
|
||||
SessionState.REJECTED:
|
||||
frozenset(),
|
||||
}
|
||||
|
||||
|
||||
class InvalidSessionTransition(RuntimeError):
|
||||
"""Raised when a session is asked to transition along an illegal edge."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Session:
|
||||
id: str = field(default_factory=lambda: uuid.uuid4().hex)
|
||||
state: SessionState = SessionState.INITIALIZING
|
||||
created_at: float = field(default_factory=time.monotonic)
|
||||
last_activity: float = field(default_factory=time.monotonic)
|
||||
|
||||
client_id: str | None = None
|
||||
preset: str | None = None
|
||||
preset_label: str | None = None
|
||||
|
||||
curated_prompts: list[str] = field(default_factory=list)
|
||||
|
||||
segment_idx: int = 0
|
||||
|
||||
enhancement_enabled: bool = False
|
||||
auto_extension_enabled: bool = False
|
||||
loop_generation_enabled: bool = False
|
||||
single_clip_mode: bool = False
|
||||
generation_paused: bool = False
|
||||
|
||||
stream_mode: str = "av_fmp4"
|
||||
gpu_id: int | None = None
|
||||
|
||||
continuation_state: ContinuationState | None = None
|
||||
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def transition(self, target: SessionState) -> None:
|
||||
"""Move to ``target`` if the edge is allowed.
|
||||
|
||||
Raises :class:`InvalidSessionTransition` on illegal moves. The
|
||||
self-loop on ``ACTIVE`` is legal so the server can re-assert
|
||||
ACTIVE on segment completion without special casing.
|
||||
"""
|
||||
allowed = _VALID_TRANSITIONS.get(self.state, frozenset())
|
||||
if target not in allowed and target is not self.state:
|
||||
raise InvalidSessionTransition(f"{self.state.value} -> {target.value} is not a valid "
|
||||
f"session transition")
|
||||
self.state = target
|
||||
self.last_activity = time.monotonic()
|
||||
|
||||
def touch(self) -> None:
|
||||
self.last_activity = time.monotonic()
|
||||
|
||||
def is_active(self) -> bool:
|
||||
return self.state is SessionState.ACTIVE
|
||||
|
||||
def segment_cap_reached(self, cap: int) -> bool:
|
||||
return self.segment_idx >= cap
|
||||
|
||||
|
||||
class SessionManager:
|
||||
"""Registers sessions and enforces per-server session limits."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
segment_cap: int,
|
||||
session_timeout_seconds: int,
|
||||
max_sessions: int = 1,
|
||||
) -> None:
|
||||
self._segment_cap = segment_cap
|
||||
self._session_timeout_seconds = session_timeout_seconds
|
||||
self._max_sessions = max_sessions
|
||||
self._sessions: dict[str, Session] = {}
|
||||
|
||||
@property
|
||||
def segment_cap(self) -> int:
|
||||
return self._segment_cap
|
||||
|
||||
@property
|
||||
def session_timeout_seconds(self) -> int:
|
||||
return self._session_timeout_seconds
|
||||
|
||||
def create(self) -> Session:
|
||||
if len(self._sessions) >= self._max_sessions:
|
||||
raise SessionRejected(f"max sessions reached ({self._max_sessions})")
|
||||
session = Session()
|
||||
self._sessions[session.id] = session
|
||||
return session
|
||||
|
||||
def get(self, session_id: str) -> Session | None:
|
||||
return self._sessions.get(session_id)
|
||||
|
||||
def close(self, session_id: str) -> None:
|
||||
self._sessions.pop(session_id, None)
|
||||
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
return session_id in self._sessions
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._sessions)
|
||||
|
||||
def active_sessions(self) -> list[Session]:
|
||||
return [s for s in self._sessions.values() if s.is_active()]
|
||||
|
||||
def reap_timed_out(self, now: float | None = None) -> list[str]:
|
||||
"""Return the ids of sessions that have exceeded the idle timeout.
|
||||
|
||||
The caller is responsible for actually closing them — this
|
||||
method only *identifies* dead sessions so the server can emit
|
||||
``session_timeout`` frames before dropping the WebSocket.
|
||||
|
||||
TODO: unused until a background driver calls it. Per-connection
|
||||
idle enforcement currently happens via asyncio.wait_for on
|
||||
receive_json; this helper catches sessions stuck before any
|
||||
receive (e.g. future QUEUED state) and is expected to be wired
|
||||
into the GPU-pool reaper.
|
||||
"""
|
||||
now = now if now is not None else time.monotonic()
|
||||
dead: list[str] = []
|
||||
for sid, session in self._sessions.items():
|
||||
if session.state in {
|
||||
SessionState.COMPLETE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
SessionState.REJECTED,
|
||||
}:
|
||||
continue
|
||||
if now - session.last_activity > self._session_timeout_seconds:
|
||||
dead.append(sid)
|
||||
return dead
|
||||
|
||||
|
||||
class SessionRejected(RuntimeError):
|
||||
"""Raised when session creation fails (queue full, auth, etc.)."""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InvalidSessionTransition",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionRejected",
|
||||
"SessionState",
|
||||
]
|
||||
@@ -1,103 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Persist the initial-image blob attached to a streaming session."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import contextlib
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
_ACCEPTED_MIMES = {
|
||||
"image/png": ".png",
|
||||
"image/jpeg": ".jpg",
|
||||
"image/jpg": ".jpg",
|
||||
"image/webp": ".webp",
|
||||
}
|
||||
|
||||
_MAX_IMAGE_BYTES = 32 * 1024 * 1024 # 32 MiB cap
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SessionInitImage:
|
||||
"""Location of the persisted init image.
|
||||
|
||||
Callers pass ``path`` to ``InputConfig.image_path``; ``display_name``
|
||||
is only used for logs.
|
||||
"""
|
||||
|
||||
path: str
|
||||
display_name: str
|
||||
mime: str
|
||||
|
||||
|
||||
def persist_session_init_image(
|
||||
payload: Any,
|
||||
*,
|
||||
output_dir: str | None = None,
|
||||
) -> SessionInitImage | None:
|
||||
"""Decode a client init-image blob and persist it to disk.
|
||||
|
||||
``payload`` shape (matches the internal UI protocol)::
|
||||
|
||||
{
|
||||
"mime": "image/png",
|
||||
"name": "ref.png",
|
||||
"data": "<base64 bytes>",
|
||||
}
|
||||
|
||||
Returns ``None`` when ``payload`` is falsy (no init image). Raises
|
||||
:class:`ValueError` on schema / size / decode errors so the caller
|
||||
can surface a user-facing ``error`` frame.
|
||||
"""
|
||||
if not payload:
|
||||
return None
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("session init image must be an object")
|
||||
|
||||
mime = payload.get("mime")
|
||||
if mime not in _ACCEPTED_MIMES:
|
||||
raise ValueError(f"session init image mime {mime!r} is not one of "
|
||||
f"{sorted(_ACCEPTED_MIMES)}")
|
||||
data_b64 = payload.get("data")
|
||||
if not isinstance(data_b64, str):
|
||||
raise ValueError("session init image data must be a base64 string")
|
||||
try:
|
||||
data = base64.b64decode(data_b64, validate=True)
|
||||
except (binascii.Error, ValueError) as exc:
|
||||
raise ValueError(f"session init image data is not valid base64: {exc}") from exc
|
||||
if len(data) > _MAX_IMAGE_BYTES:
|
||||
raise ValueError(f"session init image is {len(data)} bytes; limit is "
|
||||
f"{_MAX_IMAGE_BYTES}")
|
||||
if len(data) == 0:
|
||||
raise ValueError("session init image data is empty")
|
||||
|
||||
ext = _ACCEPTED_MIMES[mime]
|
||||
display_name = _sanitize_display_name(payload.get("name")) or f"init{ext}"
|
||||
fd, path = tempfile.mkstemp(prefix="fastvideo-init-", suffix=ext, dir=output_dir)
|
||||
try:
|
||||
with os.fdopen(fd, "wb") as f:
|
||||
f.write(data)
|
||||
except Exception:
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(path)
|
||||
raise
|
||||
return SessionInitImage(path=path, display_name=display_name, mime=mime)
|
||||
|
||||
|
||||
def _sanitize_display_name(name: Any) -> str | None:
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
name = name.strip()
|
||||
if not name:
|
||||
return None
|
||||
# Strip any path components — we only keep the leaf for logging.
|
||||
return os.path.basename(name)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"SessionInitImage",
|
||||
"persist_session_init_image",
|
||||
]
|
||||
@@ -1,206 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Session state store for the FastVideo streaming server.
|
||||
|
||||
The streaming server keeps continuation state (decoded frames + audio
|
||||
latents from the previous segment) server-side so the client doesn't
|
||||
re-upload multi-megabyte tensors each WebSocket message. Two operations
|
||||
are needed:
|
||||
|
||||
* ``snapshot(session_id) -> ContinuationState`` — serialize the current
|
||||
state so it can be exported (e.g. over HTTP) or migrated to a
|
||||
different server.
|
||||
* ``hydrate(state) -> session_id`` — load a previously serialized state
|
||||
into a new session (for resume-after-disconnect flows).
|
||||
|
||||
The store is an ABC with an :class:`InMemorySessionStore` default; Redis
|
||||
or other backends can drop in without touching the pipeline.
|
||||
|
||||
Large tensor payloads (video frames, audio latents) are kept out of the
|
||||
JSON payload via an accompanying :class:`BlobStore`. Both stores share a
|
||||
process today; they are separate types so that a future implementation
|
||||
can put blobs on S3 while keeping session metadata in Redis.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
class BlobStore(ABC):
|
||||
"""Opaque byte-blob storage keyed by id.
|
||||
|
||||
A :class:`ContinuationState` payload can reference large tensors
|
||||
stored in a :class:`BlobStore` rather than inlining them, so the
|
||||
JSON payload stays small when the state travels over the wire.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
|
||||
"""Store ``data`` and return a blob id for later retrieval."""
|
||||
|
||||
@abstractmethod
|
||||
def get(self, blob_id: str) -> bytes:
|
||||
"""Load a previously stored blob. Raises ``KeyError`` if absent."""
|
||||
|
||||
@abstractmethod
|
||||
def drop(self, blob_id: str) -> None:
|
||||
"""Remove a blob. Missing ids are a no-op."""
|
||||
|
||||
@abstractmethod
|
||||
def __contains__(self, blob_id: str) -> bool:
|
||||
...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _BlobRecord:
|
||||
data: bytes
|
||||
mime: str
|
||||
|
||||
|
||||
class InMemoryBlobStore(BlobStore):
|
||||
"""Thread-safe in-memory :class:`BlobStore` for single-process servers.
|
||||
|
||||
No eviction policy — callers are responsible for calling
|
||||
:meth:`drop` when a blob's owning state is replaced or a session
|
||||
ends. A redis- or filesystem-backed :class:`BlobStore` should
|
||||
replace this when the streaming server lands as a real service
|
||||
(PR 7.5+).
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._blobs: dict[str, _BlobRecord] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
|
||||
blob_id = uuid.uuid4().hex
|
||||
with self._lock:
|
||||
self._blobs[blob_id] = _BlobRecord(data=data, mime=mime)
|
||||
return blob_id
|
||||
|
||||
def get(self, blob_id: str) -> bytes:
|
||||
with self._lock:
|
||||
record = self._blobs.get(blob_id)
|
||||
if record is None:
|
||||
raise KeyError(f"Unknown blob id: {blob_id}")
|
||||
return record.data
|
||||
|
||||
def drop(self, blob_id: str) -> None:
|
||||
with self._lock:
|
||||
self._blobs.pop(blob_id, None)
|
||||
|
||||
def __contains__(self, blob_id: str) -> bool:
|
||||
with self._lock:
|
||||
return blob_id in self._blobs
|
||||
|
||||
def __len__(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._blobs)
|
||||
|
||||
|
||||
class SessionStore(ABC):
|
||||
"""Keyed store for per-session continuation state.
|
||||
|
||||
Implementations own the session-id → state mapping. The streaming
|
||||
server calls :meth:`store` after each segment and :meth:`snapshot`
|
||||
when a client explicitly asks for an exportable state handle.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def store(self, session_id: str, state: ContinuationState) -> None:
|
||||
"""Persist ``state`` for ``session_id``, replacing any prior value."""
|
||||
|
||||
@abstractmethod
|
||||
def snapshot(self, session_id: str) -> ContinuationState | None:
|
||||
"""Return the current state for ``session_id`` (or ``None``)."""
|
||||
|
||||
@abstractmethod
|
||||
def hydrate(
|
||||
self,
|
||||
state: ContinuationState,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
) -> str:
|
||||
"""Install ``state`` as the starting point for a session.
|
||||
|
||||
When ``session_id`` is ``None`` the store allocates a fresh id
|
||||
(UUID4); when provided the store uses it verbatim, overwriting
|
||||
any prior state at that id.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def drop(self, session_id: str) -> None:
|
||||
"""Forget a session. Missing ids are a no-op."""
|
||||
|
||||
@abstractmethod
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def __iter__(self) -> Iterator[str]:
|
||||
...
|
||||
|
||||
|
||||
class InMemorySessionStore(SessionStore):
|
||||
"""Thread-safe in-memory :class:`SessionStore`.
|
||||
|
||||
Default implementation used by single-process deployments; a future
|
||||
Redis-backed store can be dropped in without changes to the server.
|
||||
|
||||
No eviction / TTL / bounded capacity — sessions only leave via
|
||||
:meth:`drop`. The live streaming server (PR 7.5+) is responsible
|
||||
for bounding growth and for dropping any :class:`BlobStore` blobs
|
||||
referenced by a state when that state is replaced or a session
|
||||
ends; this class does not know about blobs.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._sessions: dict[str, ContinuationState] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def store(self, session_id: str, state: ContinuationState) -> None:
|
||||
with self._lock:
|
||||
self._sessions[session_id] = state
|
||||
|
||||
def snapshot(self, session_id: str) -> ContinuationState | None:
|
||||
with self._lock:
|
||||
return self._sessions.get(session_id)
|
||||
|
||||
def hydrate(
|
||||
self,
|
||||
state: ContinuationState,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
) -> str:
|
||||
sid = session_id or uuid.uuid4().hex
|
||||
with self._lock:
|
||||
self._sessions[sid] = state
|
||||
return sid
|
||||
|
||||
def drop(self, session_id: str) -> None:
|
||||
with self._lock:
|
||||
self._sessions.pop(session_id, None)
|
||||
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
with self._lock:
|
||||
return session_id in self._sessions
|
||||
|
||||
def __iter__(self) -> Iterator[str]:
|
||||
with self._lock:
|
||||
return iter(list(self._sessions))
|
||||
|
||||
def __len__(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._sessions)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BlobStore",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"SessionStore",
|
||||
]
|
||||
@@ -1,213 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""fMP4 stream encoder used by the streaming server.
|
||||
|
||||
The client's Media Source Extensions player needs a continuous fMP4
|
||||
byte stream: first an *initialization segment* (``ftyp`` + ``moov``),
|
||||
then one or more *media segments* (``moof`` + ``mdat``). We pipe raw
|
||||
RGB frames into an ffmpeg subprocess configured for fragmented output
|
||||
via ``-movflags empty_moov+default_base_moof+frag_keyframe+faststart``
|
||||
and stream the bytes back out.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import subprocess
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import numpy as np
|
||||
|
||||
|
||||
@dataclass
|
||||
class FragmentedMP4Chunk:
|
||||
"""A single fMP4 byte chunk emitted by :class:`FragmentedMP4Encoder`.
|
||||
|
||||
``kind`` identifies whether the chunk is the init segment (must be
|
||||
fed into the client's ``SourceBuffer`` first) or a media fragment.
|
||||
"""
|
||||
|
||||
kind: Literal["init", "media"]
|
||||
data: bytes
|
||||
stream_id: str
|
||||
segment_idx: int
|
||||
|
||||
|
||||
class FragmentedMP4Encoder:
|
||||
"""Stream RGB frames in, fMP4 chunks out.
|
||||
|
||||
One encoder covers one segment. The server creates a new encoder
|
||||
per :class:`ltx2_segment_start`` boundary so each segment becomes
|
||||
one media fragment the client can append independently.
|
||||
|
||||
Example::
|
||||
|
||||
encoder = FragmentedMP4Encoder(width=1024, height=576, fps=24,
|
||||
segment_idx=0)
|
||||
async with encoder:
|
||||
async for chunk in encoder.encode(frames):
|
||||
await websocket.send_bytes(chunk.data)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
width: int,
|
||||
height: int,
|
||||
fps: int,
|
||||
segment_idx: int,
|
||||
stream_id: str | None = None,
|
||||
ffmpeg_path: str = "ffmpeg",
|
||||
preset: str = "ultrafast",
|
||||
pixel_format_out: str = "yuv420p",
|
||||
extra_args: list[str] | None = None,
|
||||
) -> None:
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.fps = fps
|
||||
self.segment_idx = segment_idx
|
||||
self.stream_id = stream_id or uuid.uuid4().hex
|
||||
self._ffmpeg_path = ffmpeg_path
|
||||
self._preset = preset
|
||||
self._pixel_format_out = pixel_format_out
|
||||
self._extra_args = list(extra_args or [])
|
||||
self._proc: subprocess.Popen | None = None
|
||||
self._init_emitted = False
|
||||
|
||||
async def __aenter__(self) -> FragmentedMP4Encoder:
|
||||
self._spawn()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None:
|
||||
await self.close()
|
||||
|
||||
def _spawn(self) -> None:
|
||||
args = [
|
||||
self._ffmpeg_path,
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"error",
|
||||
"-f",
|
||||
"rawvideo",
|
||||
"-pix_fmt",
|
||||
"rgb24",
|
||||
"-s",
|
||||
f"{self.width}x{self.height}",
|
||||
"-r",
|
||||
str(self.fps),
|
||||
"-i",
|
||||
"-",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
self._preset,
|
||||
"-tune",
|
||||
"zerolatency",
|
||||
"-pix_fmt",
|
||||
self._pixel_format_out,
|
||||
"-movflags",
|
||||
"empty_moov+default_base_moof+frag_keyframe+faststart",
|
||||
"-f",
|
||||
"mp4",
|
||||
*self._extra_args,
|
||||
"-",
|
||||
]
|
||||
# stderr → DEVNULL: with -loglevel error on, the only thing
|
||||
# stderr would carry is unsolicited warnings. Piping without a
|
||||
# reader deadlocks ffmpeg once the pipe buffer (~64 KiB) fills.
|
||||
self._proc = subprocess.Popen( # noqa: S603
|
||||
args,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
bufsize=0,
|
||||
)
|
||||
|
||||
async def encode(
|
||||
self,
|
||||
frames: list[np.ndarray] | AsyncIterator[np.ndarray],
|
||||
) -> AsyncIterator[FragmentedMP4Chunk]:
|
||||
"""Feed frames into ffmpeg and yield fMP4 chunks as they appear."""
|
||||
if self._proc is None:
|
||||
self._spawn()
|
||||
assert self._proc is not None and self._proc.stdin is not None
|
||||
proc = self._proc
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
async def _writer() -> None:
|
||||
try:
|
||||
if hasattr(frames, "__aiter__"):
|
||||
async for frame in frames: # type: ignore[union-attr]
|
||||
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
|
||||
else:
|
||||
for frame in frames: # type: ignore[assignment]
|
||||
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
|
||||
finally:
|
||||
with contextlib.suppress(BrokenPipeError):
|
||||
proc.stdin.close()
|
||||
|
||||
writer_task = asyncio.create_task(_writer())
|
||||
try:
|
||||
reader = proc.stdout
|
||||
assert reader is not None
|
||||
# Read in reasonably-sized chunks; MSE tolerates any size
|
||||
# but we don't want to starve the event loop.
|
||||
chunk_size = 64 * 1024
|
||||
while True:
|
||||
data = await loop.run_in_executor(None, reader.read, chunk_size)
|
||||
if not data:
|
||||
break
|
||||
kind: Literal["init", "media"] = "init" if not self._init_emitted else "media"
|
||||
self._init_emitted = True
|
||||
yield FragmentedMP4Chunk(
|
||||
kind=kind,
|
||||
data=bytes(data),
|
||||
stream_id=self.stream_id,
|
||||
segment_idx=self.segment_idx,
|
||||
)
|
||||
finally:
|
||||
await writer_task
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._proc is None:
|
||||
return
|
||||
proc = self._proc
|
||||
self._proc = None
|
||||
try:
|
||||
if proc.stdin and not proc.stdin.closed:
|
||||
proc.stdin.close()
|
||||
except BrokenPipeError:
|
||||
pass
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
loop.run_in_executor(None, proc.wait),
|
||||
timeout=5.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
proc.kill()
|
||||
await loop.run_in_executor(None, proc.wait)
|
||||
|
||||
|
||||
def _write_frame(stdin, frame: np.ndarray) -> None:
|
||||
import numpy as np
|
||||
|
||||
if not isinstance(frame, np.ndarray):
|
||||
raise TypeError("fMP4 encoder frames must be numpy.ndarray")
|
||||
if frame.dtype != np.uint8:
|
||||
frame = frame.astype(np.uint8)
|
||||
if frame.ndim != 3 or frame.shape[-1] != 3:
|
||||
raise ValueError("fMP4 encoder frames must be HxWx3 uint8 RGB; got "
|
||||
f"shape={frame.shape}, dtype={frame.dtype}")
|
||||
with contextlib.suppress(BrokenPipeError):
|
||||
stdin.write(frame.tobytes())
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
]
|
||||
@@ -627,32 +627,18 @@ class VideoGenerator:
|
||||
gen_time = time.perf_counter() - start_time
|
||||
logger.info("Generated successfully in %.2f seconds", gen_time)
|
||||
|
||||
# Process outputs (skip the make_grid loop for audio-only, where
|
||||
# `samples` is a 1×3×1×8×8 placeholder no caller will use).
|
||||
audio_only = bool(output_batch.extra.get("audio_only"))
|
||||
frames: list[np.ndarray] = []
|
||||
if not audio_only:
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.permute(1, 2, 0).squeeze(-1)
|
||||
x = (x * 255).to(torch.uint8)
|
||||
frames.append(x.cpu().numpy())
|
||||
# Process outputs
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.permute(1, 2, 0).squeeze(-1)
|
||||
x = (x * 255).to(torch.uint8)
|
||||
frames.append(x.cpu().numpy())
|
||||
|
||||
# Save output if requested
|
||||
if batch.save_video:
|
||||
if output_batch.extra.get("audio_only"):
|
||||
# Audio-only workload: write a standalone .wav rather than
|
||||
# muxing the audio into a placeholder mp4 (which forces
|
||||
# ffmpeg to round 8x8 placeholder frames up to 16x16).
|
||||
output_path = self._rewrite_extension(output_path, ".wav")
|
||||
self._write_pcm_wav(
|
||||
output_path,
|
||||
output_batch.extra["audio"],
|
||||
int(output_batch.extra["audio_sample_rate"]),
|
||||
)
|
||||
logger.info("Saved audio to %s", output_path)
|
||||
elif self._is_image_workload():
|
||||
if self._is_image_workload():
|
||||
# Image workloads (t2i, i2i, …): save the first frame as PNG.
|
||||
imageio.imwrite(output_path, frames[0])
|
||||
logger.info("Saved image to %s", output_path)
|
||||
@@ -669,11 +655,7 @@ class VideoGenerator:
|
||||
"prompts": prompt,
|
||||
"samples": samples if batch.return_frames else None,
|
||||
"frames": frames if batch.return_frames else None,
|
||||
# Audio is the primary output for audio workloads — return it
|
||||
# whenever the pipeline produced one, regardless of
|
||||
# `return_frames` (which gates the video-shaped buffers).
|
||||
"audio": output_batch.extra.get("audio"),
|
||||
"audio_sample_rate": output_batch.extra.get("audio_sample_rate"),
|
||||
"audio": output_batch.extra.get("audio") if batch.return_frames else None,
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
"logging_info": logging_info,
|
||||
@@ -701,55 +683,7 @@ class VideoGenerator:
|
||||
return result.to_legacy_dict()
|
||||
|
||||
@staticmethod
|
||||
def _rewrite_extension(path: str, new_ext: str) -> str:
|
||||
root, old_ext = os.path.splitext(path)
|
||||
new_path = root + new_ext
|
||||
if old_ext and old_ext.lower() != new_ext.lower():
|
||||
logger.info("Rewriting output extension %s -> %s.", old_ext, new_ext)
|
||||
return new_path
|
||||
|
||||
@staticmethod
|
||||
def _audio_to_int16(audio: torch.Tensor | np.ndarray, ) -> tuple[np.ndarray, int]:
|
||||
"""Normalize `[samples]` / `[samples, channels]` / `[channels,
|
||||
samples]` audio in roughly [-1, 1] to a `(int16 [samples,
|
||||
channels], num_channels)` pair. Raises `ValueError` for shapes
|
||||
we can't classify.
|
||||
"""
|
||||
if torch.is_tensor(audio):
|
||||
audio_np = audio.detach().cpu().float().numpy()
|
||||
else:
|
||||
audio_np = np.asarray(audio, dtype=np.float32)
|
||||
if audio_np.ndim == 1:
|
||||
audio_np = audio_np[:, None]
|
||||
elif audio_np.ndim == 2:
|
||||
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
|
||||
audio_np = audio_np.T
|
||||
else:
|
||||
raise ValueError(f"Unexpected audio shape {audio_np.shape}.")
|
||||
audio_np = np.clip(audio_np, -1.0, 1.0)
|
||||
audio_int16 = (audio_np * 32767.0).astype(np.int16)
|
||||
return audio_int16, audio_int16.shape[1]
|
||||
|
||||
@classmethod
|
||||
def _write_pcm_wav(
|
||||
cls,
|
||||
wav_path: str,
|
||||
audio: torch.Tensor | np.ndarray,
|
||||
sample_rate: int,
|
||||
) -> int:
|
||||
"""Write 16-bit PCM WAV; returns the channel count."""
|
||||
import wave
|
||||
audio_int16, num_channels = cls._audio_to_int16(audio)
|
||||
with wave.open(wav_path, "wb") as f:
|
||||
f.setnchannels(num_channels)
|
||||
f.setsampwidth(2)
|
||||
f.setframerate(sample_rate)
|
||||
f.writeframes(audio_int16.tobytes())
|
||||
return num_channels
|
||||
|
||||
@classmethod
|
||||
def _mux_audio(
|
||||
cls,
|
||||
video_path: str,
|
||||
audio: torch.Tensor | np.ndarray,
|
||||
sample_rate: int,
|
||||
@@ -762,13 +696,37 @@ class VideoGenerator:
|
||||
"Install with: pip install av")
|
||||
return False
|
||||
|
||||
if torch.is_tensor(audio):
|
||||
audio_np = audio.detach().cpu().float().numpy()
|
||||
else:
|
||||
audio_np = np.asarray(audio, dtype=np.float32)
|
||||
|
||||
if audio_np.ndim == 1:
|
||||
audio_np = audio_np[:, None]
|
||||
elif audio_np.ndim == 2:
|
||||
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
|
||||
audio_np = audio_np.T
|
||||
else:
|
||||
logger.warning("Unexpected audio shape %s; skipping mux.", audio_np.shape)
|
||||
return False
|
||||
|
||||
audio_np = np.clip(audio_np, -1.0, 1.0)
|
||||
audio_int16 = (audio_np * 32767.0).astype(np.int16)
|
||||
num_channels = audio_int16.shape[1]
|
||||
layout = "stereo" if num_channels == 2 else "mono"
|
||||
|
||||
try:
|
||||
import wave
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
out_path = os.path.join(tmpdir, "muxed.mp4")
|
||||
wav_path = os.path.join(tmpdir, "audio.wav")
|
||||
|
||||
num_channels = cls._write_pcm_wav(wav_path, audio, sample_rate)
|
||||
layout = "stereo" if num_channels == 2 else "mono"
|
||||
# Write audio to WAV file
|
||||
with wave.open(wav_path, "wb") as wav_file:
|
||||
wav_file.setnchannels(num_channels)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(sample_rate)
|
||||
wav_file.writeframes(audio_int16.tobytes())
|
||||
|
||||
# Open input video and audio
|
||||
input_video = av.open(video_path)
|
||||
|
||||
@@ -844,13 +844,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
|
||||
min_timestep_ratio: float = 0.2
|
||||
max_timestep_ratio: float = 0.98
|
||||
# CFG scale applied to the real (teacher) score in the DMD loss, using the
|
||||
# parameterization `x = x_cond + w * (x_cond - x_uncond)`. This differs
|
||||
# from the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)` by an
|
||||
# offset of 1: `w_here = w_standard - 1`. So `w=0` recovers the
|
||||
# conditional output, `w=-1` recovers the unconditional output, and the
|
||||
# default 3.5 corresponds to a standard CFG scale of 4.5. Matches the
|
||||
# original DMD2 reference implementation.
|
||||
real_score_guidance_scale: float = 3.5
|
||||
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
|
||||
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
|
||||
@@ -1111,10 +1104,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--real-score-guidance-scale",
|
||||
type=float,
|
||||
default=TrainingArgs.real_score_guidance_scale,
|
||||
help=("Teacher CFG scale for the real score in the DMD loss. Uses "
|
||||
"the parameterization x_cond + w * (x_cond - x_uncond), so "
|
||||
"w=0 -> cond, w=-1 -> uncond, and the relation to standard "
|
||||
"CFG is w_standard = w + 1 (default 3.5 == standard 4.5)."))
|
||||
help="Teacher guidance scale")
|
||||
parser.add_argument("--fake-score-learning-rate",
|
||||
type=float,
|
||||
default=TrainingArgs.fake_score_learning_rate,
|
||||
|
||||
@@ -556,10 +556,7 @@ class CosmosTransformer3DModel(BaseDiT):
|
||||
self.extra_pos_embed_type = config.extra_pos_embed_type
|
||||
|
||||
# 1. Patch Embedding
|
||||
# config.in_channels already includes the condition_mask channel
|
||||
# (HF config: in_channels=17 = 16 latent + 1 condition_mask).
|
||||
# Only add +1 for the padding_mask when concat_padding_mask=True.
|
||||
patch_embed_in_channels = config.in_channels + (1 if config.concat_padding_mask else 0)
|
||||
patch_embed_in_channels = config.in_channels + 1 if config.concat_padding_mask else config.in_channels
|
||||
self.patch_embed = CosmosPatchEmbed(patch_embed_in_channels,
|
||||
inner_dim,
|
||||
config.patch_size,
|
||||
@@ -620,28 +617,6 @@ class CosmosTransformer3DModel(BaseDiT):
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
|
||||
# Defensive dtype alignment: the Cosmos checkpoint is bf16 but
|
||||
# FSDP-wrapped training copies may report fp32 via
|
||||
# `next(parameters()).dtype`, which disables autocast in the
|
||||
# shared denoising stage and feeds fp32 tensors into bf16
|
||||
# weights. Cast every external input to the patch_embed weight
|
||||
# dtype so the model forward is robust regardless of caller.
|
||||
_target_dtype = self.patch_embed.proj.weight.dtype
|
||||
if hidden_states.dtype != _target_dtype:
|
||||
hidden_states = hidden_states.to(_target_dtype)
|
||||
if condition_mask is not None and condition_mask.dtype != _target_dtype:
|
||||
condition_mask = condition_mask.to(_target_dtype)
|
||||
if padding_mask is not None and padding_mask.dtype != _target_dtype:
|
||||
padding_mask = padding_mask.to(_target_dtype)
|
||||
if isinstance(encoder_hidden_states, torch.Tensor):
|
||||
if encoder_hidden_states.dtype != _target_dtype:
|
||||
encoder_hidden_states = encoder_hidden_states.to(_target_dtype)
|
||||
else:
|
||||
encoder_hidden_states = [
|
||||
t.to(_target_dtype) if t.dtype != _target_dtype else t
|
||||
for t in encoder_hidden_states
|
||||
]
|
||||
|
||||
# 1. Concatenate padding mask if needed & prepare attention mask
|
||||
if condition_mask is not None:
|
||||
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
|
||||
@@ -651,10 +626,6 @@ class CosmosTransformer3DModel(BaseDiT):
|
||||
padding_mask = transforms.functional.resize(
|
||||
padding_mask, list(hidden_states.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST
|
||||
)
|
||||
# torchvision.resize may upcast bf16/fp16 → fp32; restore
|
||||
# hidden_states' dtype so the subsequent cat doesn't promote
|
||||
# everything and break patch_embed (bf16 weights).
|
||||
padding_mask = padding_mask.to(hidden_states.dtype)
|
||||
hidden_states = torch.cat(
|
||||
[hidden_states, padding_mask.unsqueeze(2).repeat(batch_size, 1, num_frames, 1, 1)], dim=1
|
||||
)
|
||||
@@ -732,11 +703,6 @@ class CosmosTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5)
|
||||
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
# Return as tuple for compatibility with callers that
|
||||
# do `transformer(..., return_dict=False)[0]` (diffusers
|
||||
# convention used by CosmosDenoisingStage).
|
||||
if not kwargs.get("return_dict", True):
|
||||
return (hidden_states,)
|
||||
return hidden_states
|
||||
|
||||
# Entry point for model registry
|
||||
|
||||
@@ -1,389 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 DiT.
|
||||
|
||||
Continuous transformer with rotary self-attention, GQA cross-attention,
|
||||
and prepend global conditioning. 24 layers, embed_dim=1536, head_dim=64.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.configs.models.dits import StableAudioConfig
|
||||
from fastvideo.layers.layernorm import FP32LayerNorm
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.loader.utils import get_param_names_mapping
|
||||
|
||||
# Single import-time snapshot — re-reading via `StableAudioConfig()` per
|
||||
# `Attention.__init__` would rebuild the nested dataclass + regex map ~48
|
||||
# times during a single DiT construction. Reused for the class-level
|
||||
# attribute defaults below.
|
||||
_DEFAULT_CONFIG = StableAudioConfig()
|
||||
_SUPPORTED_BACKENDS = _DEFAULT_CONFIG.arch_config._supported_attention_backends
|
||||
|
||||
|
||||
class FourierFeatures(nn.Module):
|
||||
"""Random-Fourier learned-frequency timestep encoder."""
|
||||
|
||||
def __init__(self, in_features: int, out_features: int, std: float = 1.0) -> None:
|
||||
super().__init__()
|
||||
assert out_features % 2 == 0
|
||||
self.weight = nn.Parameter(torch.randn([out_features // 2, in_features]) * std)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
f = 2 * math.pi * x @ self.weight.T
|
||||
return torch.cat([f.cos(), f.sin()], dim=-1)
|
||||
|
||||
|
||||
# Partial-rotary with halves-swap (`unbind(-2)`, `[-x2, x1]`). Different
|
||||
# from FastVideo's `_apply_rotary_emb` (interleaved pairs, `unbind(-1)`),
|
||||
# so kept local.
|
||||
|
||||
|
||||
class RotaryEmbedding(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, base: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
inv_freq = 1.0 / (base**(torch.arange(0, dim, 2).float() / dim))
|
||||
self.register_buffer("inv_freq", inv_freq)
|
||||
self.register_buffer("scale", None)
|
||||
|
||||
def forward_from_seq_len(self, seq_len: int):
|
||||
t = torch.arange(seq_len, device=self.inv_freq.device, dtype=torch.float32)
|
||||
freqs = torch.einsum("i , j -> i j", t, self.inv_freq)
|
||||
freqs = torch.cat((freqs, freqs), dim=-1)
|
||||
return freqs, 1.0
|
||||
|
||||
|
||||
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(x, "... (j d) -> ... j d", j=2)
|
||||
x1, x2 = x.unbind(dim=-2)
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
|
||||
def _apply_rotary_pos_emb(t: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
|
||||
out_dtype = t.dtype
|
||||
rot_dim, seq_len = freqs.shape[-1], t.shape[-2]
|
||||
freqs = freqs.to(torch.float32)[-seq_len:, :]
|
||||
t = t.to(torch.float32)
|
||||
if t.ndim == 4 and freqs.ndim == 3:
|
||||
freqs = rearrange(freqs, "b n d -> b 1 n d")
|
||||
t_rot, t_unrot = t[..., :rot_dim], t[..., rot_dim:]
|
||||
t_rot = (t_rot * freqs.cos()) + (_rotate_half(t_rot) * freqs.sin())
|
||||
return torch.cat((t_rot.to(out_dtype), t_unrot.to(out_dtype)), dim=-1)
|
||||
|
||||
|
||||
# SwiGLU FF — local because `fastvideo.layers.mlp.MLP` is non-gated.
|
||||
|
||||
|
||||
class _GLU(nn.Module):
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, activation: nn.Module) -> None:
|
||||
super().__init__()
|
||||
self.act = activation
|
||||
self.proj = ReplicatedLinear(dim_in, dim_out * 2, bias=True)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x, _ = self.proj(x)
|
||||
x, gate = x.chunk(2, dim=-1)
|
||||
return x * self.act(gate)
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
# Sequential layout `(GLU, Identity, Linear, Identity)` keeps the
|
||||
# checkpoint keys at indices 0 and 2.
|
||||
|
||||
def __init__(self, dim: int, mult: int = 4, zero_init_output: bool = True) -> None:
|
||||
super().__init__()
|
||||
inner_dim = int(dim * mult)
|
||||
linear_in = _GLU(dim, inner_dim, nn.SiLU())
|
||||
linear_out = ReplicatedLinear(inner_dim, dim, bias=True)
|
||||
if zero_init_output:
|
||||
nn.init.zeros_(linear_out.weight)
|
||||
nn.init.zeros_(linear_out.bias)
|
||||
self.ff = nn.Sequential(linear_in, nn.Identity(), linear_out, nn.Identity())
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
for mod in self.ff:
|
||||
if isinstance(mod, ReplicatedLinear):
|
||||
x, _ = mod(x)
|
||||
else:
|
||||
x = mod(x)
|
||||
return x
|
||||
|
||||
|
||||
# Cross-attention is GQA (24 query heads, 12 KV heads); both backends
|
||||
# (FlashAttn, SDPA with `enable_gqa=True`) handle it.
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, dim_heads: int = 64, dim_context: int | None = None,
|
||||
zero_init_output: bool = True, qk_norm: str | None = None) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.dim_heads = dim_heads
|
||||
dim_kv = dim_context if dim_context is not None else dim
|
||||
self.num_heads = dim // dim_heads
|
||||
self.kv_heads = dim_kv // dim_heads
|
||||
if dim_context is not None:
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=False)
|
||||
self.to_kv = ReplicatedLinear(dim_kv, dim_kv * 2, bias=False)
|
||||
else:
|
||||
self.to_qkv = ReplicatedLinear(dim, dim * 3, bias=False)
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=False)
|
||||
if zero_init_output:
|
||||
nn.init.zeros_(self.to_out.weight)
|
||||
|
||||
# `stable-audio-open-small` wraps Q/K in LayerNorm before attn
|
||||
# (`attn_kwargs.qk_norm = "ln"` in its `model_config.json`); the
|
||||
# 1.0 base does not. Names match upstream (`q_norm`/`k_norm`)
|
||||
# so the converted state dict loads strict.
|
||||
if qk_norm == "ln":
|
||||
self.q_norm = nn.LayerNorm(dim_heads)
|
||||
self.k_norm = nn.LayerNorm(dim_heads)
|
||||
elif qk_norm is None:
|
||||
self.q_norm = nn.Identity()
|
||||
self.k_norm = nn.Identity()
|
||||
else:
|
||||
raise ValueError(f"Unsupported qk_norm={qk_norm!r}; expected 'ln' or None.")
|
||||
|
||||
self.attn = LocalAttention(num_heads=self.num_heads, head_size=dim_heads,
|
||||
num_kv_heads=self.kv_heads, causal=False,
|
||||
supported_attention_backends=_SUPPORTED_BACKENDS)
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor | None = None,
|
||||
rotary_pos_emb: tuple[torch.Tensor, float] | None = None) -> torch.Tensor:
|
||||
h, kv_h, has_context = self.num_heads, self.kv_heads, context is not None
|
||||
kv_input = context if has_context else x
|
||||
if has_context:
|
||||
q, _ = self.to_q(x)
|
||||
kv, _ = self.to_kv(kv_input)
|
||||
k, v = kv.chunk(2, dim=-1)
|
||||
else:
|
||||
qkv, _ = self.to_qkv(x)
|
||||
q, k, v = qkv.chunk(3, dim=-1)
|
||||
# LocalAttention expects [batch, seq_len, num_heads, head_dim].
|
||||
q = rearrange(q, "b n (h d) -> b n h d", h=h)
|
||||
k = rearrange(k, "b n (h d) -> b n h d", h=kv_h)
|
||||
v = rearrange(v, "b n (h d) -> b n h d", h=kv_h)
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
if rotary_pos_emb is not None:
|
||||
freqs, _ = rotary_pos_emb
|
||||
v_dtype = v.dtype
|
||||
# Partial rotary (rot_dim < head_dim) with halves-swap, so
|
||||
# apply outside LocalAttention. q,k come in as [B, S, H, D];
|
||||
# transpose to [B, H, S, D] for the helper.
|
||||
q_t = q.transpose(1, 2)
|
||||
k_t = k.transpose(1, 2)
|
||||
if q_t.shape[-2] >= k_t.shape[-2]:
|
||||
ratio = q_t.shape[-2] / k_t.shape[-2]
|
||||
q_freqs, k_freqs = freqs, ratio * freqs
|
||||
else:
|
||||
ratio = k_t.shape[-2] / q_t.shape[-2]
|
||||
q_freqs, k_freqs = ratio * freqs, freqs
|
||||
q = _apply_rotary_pos_emb(q_t, q_freqs).to(v_dtype).transpose(1, 2)
|
||||
k = _apply_rotary_pos_emb(k_t, k_freqs).to(v_dtype).transpose(1, 2)
|
||||
|
||||
out = self.attn(q, k, v)
|
||||
out = rearrange(out, "b n h d -> b n (h d)")
|
||||
out, _ = self.to_out(out)
|
||||
return out
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, dim_heads: int = 64, cross_attend: bool = False,
|
||||
dim_context: int | None = None, zero_init_branch_outputs: bool = True,
|
||||
qk_norm: str | None = None) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.dim_heads = min(dim_heads, dim)
|
||||
self.cross_attend = cross_attend
|
||||
self.pre_norm = FP32LayerNorm(dim, elementwise_affine=True)
|
||||
self.self_attn = Attention(dim, dim_heads=self.dim_heads,
|
||||
zero_init_output=zero_init_branch_outputs,
|
||||
qk_norm=qk_norm)
|
||||
if cross_attend:
|
||||
self.cross_attend_norm = FP32LayerNorm(dim, elementwise_affine=True)
|
||||
self.cross_attn = Attention(dim, dim_heads=self.dim_heads, dim_context=dim_context,
|
||||
zero_init_output=zero_init_branch_outputs,
|
||||
qk_norm=qk_norm)
|
||||
self.ff_norm = FP32LayerNorm(dim, elementwise_affine=True)
|
||||
self.ff = FeedForward(dim, zero_init_output=zero_init_branch_outputs)
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor | None = None,
|
||||
rotary_pos_emb: tuple[torch.Tensor, float] | None = None) -> torch.Tensor:
|
||||
x = x + self.self_attn(self.pre_norm(x), rotary_pos_emb=rotary_pos_emb)
|
||||
if context is not None and self.cross_attend:
|
||||
x = x + self.cross_attn(self.cross_attend_norm(x), context=context)
|
||||
x = x + self.ff(self.ff_norm(x))
|
||||
return x
|
||||
|
||||
|
||||
class ContinuousTransformer(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, depth: int, *, dim_heads: int = 64, dim_in: int | None = None,
|
||||
dim_out: int | None = None, cross_attend: bool = False,
|
||||
cond_token_dim: int | None = None, zero_init_branch_outputs: bool = True,
|
||||
qk_norm: str | None = None) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.depth = depth
|
||||
self.project_in = (ReplicatedLinear(dim_in, dim, bias=False) if dim_in is not None
|
||||
else nn.Identity())
|
||||
self.project_out = (ReplicatedLinear(dim, dim_out, bias=False) if dim_out is not None
|
||||
else nn.Identity())
|
||||
self.rotary_pos_emb = RotaryEmbedding(max(dim_heads // 2, 32))
|
||||
self.layers = nn.ModuleList([
|
||||
TransformerBlock(dim, dim_heads=dim_heads, cross_attend=cross_attend,
|
||||
dim_context=cond_token_dim,
|
||||
zero_init_branch_outputs=zero_init_branch_outputs,
|
||||
qk_norm=qk_norm) for _ in range(depth)
|
||||
])
|
||||
|
||||
def forward(self, x: torch.Tensor, prepend_embeds: torch.Tensor | None = None,
|
||||
context: torch.Tensor | None = None) -> torch.Tensor:
|
||||
if isinstance(self.project_in, ReplicatedLinear):
|
||||
x, _ = self.project_in(x)
|
||||
if prepend_embeds is not None:
|
||||
assert prepend_embeds.shape[-1] == x.shape[-1]
|
||||
x = torch.cat((prepend_embeds, x), dim=-2)
|
||||
rotary = self.rotary_pos_emb.forward_from_seq_len(x.shape[1])
|
||||
for layer in self.layers:
|
||||
x = layer(x, context=context, rotary_pos_emb=rotary)
|
||||
if isinstance(self.project_out, ReplicatedLinear):
|
||||
x, _ = self.project_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class StableAudioDiT(BaseDiT):
|
||||
"""Stable Audio Open 1.0 diffusion transformer."""
|
||||
|
||||
_fsdp_shard_conditions = _DEFAULT_CONFIG.arch_config._fsdp_shard_conditions
|
||||
_compile_conditions = _DEFAULT_CONFIG.arch_config._compile_conditions
|
||||
param_names_mapping = _DEFAULT_CONFIG.arch_config.param_names_mapping
|
||||
reverse_param_names_mapping: dict = {}
|
||||
|
||||
def __init__(self, config: StableAudioConfig | None = None,
|
||||
hf_config: dict[str, Any] | None = None) -> None:
|
||||
if config is None:
|
||||
config = StableAudioConfig()
|
||||
super().__init__(config=config, hf_config=hf_config or {})
|
||||
arch = config.arch_config
|
||||
self.hidden_size = arch.hidden_size
|
||||
self.num_attention_heads = arch.num_attention_heads
|
||||
self.num_channels_latents = arch.num_channels_latents
|
||||
io_channels = arch.io_channels
|
||||
embed_dim = arch.embed_dim
|
||||
depth = arch.depth
|
||||
num_heads = arch.num_attention_heads
|
||||
cond_token_dim = arch.cond_token_dim
|
||||
global_cond_dim = arch.global_cond_dim
|
||||
project_cond_tokens = arch.project_cond_tokens
|
||||
project_global_cond = arch.project_global_cond
|
||||
qk_norm = arch.qk_norm
|
||||
|
||||
self.cond_token_dim = cond_token_dim
|
||||
timestep_features_dim = 256
|
||||
self.timestep_features = FourierFeatures(1, timestep_features_dim)
|
||||
self.to_timestep_embed = nn.Sequential(
|
||||
ReplicatedLinear(timestep_features_dim, embed_dim, bias=True),
|
||||
nn.SiLU(),
|
||||
ReplicatedLinear(embed_dim, embed_dim, bias=True),
|
||||
)
|
||||
self.diffusion_objective = "v"
|
||||
|
||||
cond_embed_dim = cond_token_dim if not project_cond_tokens else embed_dim
|
||||
self.to_cond_embed = nn.Sequential(
|
||||
ReplicatedLinear(cond_token_dim, cond_embed_dim, bias=False),
|
||||
nn.SiLU(),
|
||||
ReplicatedLinear(cond_embed_dim, cond_embed_dim, bias=False),
|
||||
)
|
||||
|
||||
global_embed_dim = global_cond_dim if not project_global_cond else embed_dim
|
||||
self.to_global_embed = nn.Sequential(
|
||||
ReplicatedLinear(global_cond_dim, global_embed_dim, bias=False),
|
||||
nn.SiLU(),
|
||||
ReplicatedLinear(global_embed_dim, global_embed_dim, bias=False),
|
||||
)
|
||||
|
||||
self.transformer = ContinuousTransformer(
|
||||
dim=embed_dim, depth=depth, dim_heads=embed_dim // num_heads, dim_in=io_channels,
|
||||
dim_out=io_channels, cross_attend=True, cond_token_dim=cond_embed_dim,
|
||||
qk_norm=qk_norm,
|
||||
)
|
||||
|
||||
self.preprocess_conv = nn.Conv1d(io_channels, io_channels, 1, bias=False)
|
||||
nn.init.zeros_(self.preprocess_conv.weight)
|
||||
self.postprocess_conv = nn.Conv1d(io_channels, io_channels, 1, bias=False)
|
||||
nn.init.zeros_(self.postprocess_conv.weight)
|
||||
|
||||
self.io_channels = io_channels
|
||||
self.embed_dim = embed_dim
|
||||
self.depth = depth
|
||||
self.num_heads = num_heads
|
||||
self.__post_init__()
|
||||
|
||||
@staticmethod
|
||||
def _seq_apply(seq: nn.Sequential, x: torch.Tensor) -> torch.Tensor:
|
||||
for mod in seq:
|
||||
if isinstance(mod, ReplicatedLinear):
|
||||
x, _ = mod(x)
|
||||
else:
|
||||
x = mod(x)
|
||||
return x
|
||||
|
||||
def forward(self, x: torch.Tensor, t: torch.Tensor, *, cross_attn_cond: torch.Tensor,
|
||||
global_embed: torch.Tensor) -> torch.Tensor:
|
||||
"""Forward over a single batch. CFG batching is the caller's job."""
|
||||
model_dtype = next(self.parameters()).dtype
|
||||
x = x.to(model_dtype)
|
||||
t = t.to(model_dtype)
|
||||
cross_attn_cond = cross_attn_cond.to(model_dtype)
|
||||
global_embed = global_embed.to(model_dtype)
|
||||
|
||||
cross_attn_cond = self._seq_apply(self.to_cond_embed, cross_attn_cond)
|
||||
global_embed = self._seq_apply(self.to_global_embed, global_embed)
|
||||
timestep_embed = self._seq_apply(self.to_timestep_embed, self.timestep_features(t[:, None]))
|
||||
global_embed = global_embed + timestep_embed
|
||||
prepend_inputs = global_embed.unsqueeze(1)
|
||||
|
||||
x = self.preprocess_conv(x) + x
|
||||
x = rearrange(x, "b c t -> b t c")
|
||||
out = self.transformer(x, prepend_embeds=prepend_inputs, context=cross_attn_cond)
|
||||
out = rearrange(out, "b t c -> b c t")[:, :, prepend_inputs.shape[1]:]
|
||||
return self.postprocess_conv(out) + out
|
||||
|
||||
@classmethod
|
||||
def from_official_state_dict(cls, state_dict: dict[str, torch.Tensor],
|
||||
prefix: str = "model.model.") -> "StableAudioDiT":
|
||||
"""Load from a raw `stable_audio_tools` monolithic state dict.
|
||||
Kept for tests / older checkpoints; production loads go through
|
||||
the standard `TransformerLoader` against the converted Diffusers
|
||||
repo.
|
||||
"""
|
||||
model = cls()
|
||||
mapping_fn = get_param_names_mapping(model.config.arch_config.param_names_mapping)
|
||||
remapped: dict[str, torch.Tensor] = {}
|
||||
for k, v in state_dict.items():
|
||||
if not k.startswith(prefix):
|
||||
continue
|
||||
new_key, _, _ = mapping_fn(k)
|
||||
remapped[new_key] = v
|
||||
missing, unexpected = model.load_state_dict(remapped, strict=True)
|
||||
if missing or unexpected:
|
||||
raise RuntimeError(
|
||||
f"StableAudioDiT load mismatch — missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model
|
||||
|
||||
|
||||
EntryClass = StableAudioDiT
|
||||
@@ -1,214 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 conditioner.
|
||||
|
||||
T5-base text encoder + two NumberConditioners (`seconds_start`,
|
||||
`seconds_total`), wrapped by `StableAudioMultiConditioner` which
|
||||
produces the cross-attention and global-conditioning tensors the DiT
|
||||
expects.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.configs.models.encoders import StableAudioConditionerConfig
|
||||
|
||||
|
||||
class _LearnedPositionalEmbedding(nn.Module):
|
||||
|
||||
def __init__(self, dim: int) -> None:
|
||||
super().__init__()
|
||||
assert (dim % 2) == 0
|
||||
self.weights = nn.Parameter(torch.randn(dim // 2))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(x, "b -> b 1")
|
||||
freqs = x * rearrange(self.weights, "d -> 1 d") * 2 * math.pi
|
||||
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim=-1)
|
||||
return torch.cat((x, fouriered), dim=-1)
|
||||
|
||||
|
||||
def _time_positional_embedding(dim: int, out_features: int) -> nn.Sequential:
|
||||
return nn.Sequential(_LearnedPositionalEmbedding(dim),
|
||||
nn.Linear(in_features=dim + 1, out_features=out_features))
|
||||
|
||||
|
||||
class NumberEmbedder(nn.Module):
|
||||
|
||||
def __init__(self, features: int, dim: int = 256) -> None:
|
||||
super().__init__()
|
||||
self.features = features
|
||||
self.embedding = _time_positional_embedding(dim=dim, out_features=features)
|
||||
|
||||
def forward(self, x: torch.Tensor | list[float]) -> torch.Tensor:
|
||||
if not torch.is_tensor(x):
|
||||
device = next(self.embedding.parameters()).device
|
||||
x = torch.tensor(x, device=device)
|
||||
shape = x.shape
|
||||
x = rearrange(x, "... -> (...)")
|
||||
out = self.embedding(x)
|
||||
return out.view(*shape, self.features)
|
||||
|
||||
|
||||
class _Conditioner(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, output_dim: int, project_out: bool = False) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.output_dim = output_dim
|
||||
self.proj_out = (nn.Linear(dim, output_dim) if dim != output_dim or project_out
|
||||
else nn.Identity())
|
||||
|
||||
|
||||
class T5Conditioner(_Conditioner):
|
||||
"""T5 text conditioner. Pads to `model_max_length` (=128 for the SA
|
||||
repo's tokenizer, NOT the standard 512) and emits a masked
|
||||
last-hidden-state.
|
||||
"""
|
||||
|
||||
T5_MODEL_DIMS = {"t5-base": 768}
|
||||
|
||||
def __init__(self, output_dim: int, t5_model_name: str = "t5-base",
|
||||
max_length: int = 128, dtype: str = "float16") -> None:
|
||||
super().__init__(self.T5_MODEL_DIMS[t5_model_name], output_dim, project_out=False)
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
self.max_length = max_length
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(t5_model_name)
|
||||
# T5 loaded directly in fp16 (config-driven) to match official
|
||||
# `stable_audio_tools/models/conditioners.py:334`. Registered as
|
||||
# a normal submodule so `.to(device)` / `torch.compile` track it;
|
||||
# `from_official_state_dict` filters `conditioners.prompt.*` from
|
||||
# the missing-key check (T5 weights are absent from the SA
|
||||
# checkpoint by design).
|
||||
# Explicit lookup so a typo (e.g. "fp16" instead of "float16") errors
|
||||
# at load time rather than silently falling back to a wrong dtype.
|
||||
torch_dtype = getattr(torch, dtype)
|
||||
if not isinstance(torch_dtype, torch.dtype):
|
||||
raise ValueError(f"T5Conditioner dtype={dtype!r} is not a torch.dtype.")
|
||||
self._t5_dtype = torch_dtype
|
||||
self.model = (T5EncoderModel.from_pretrained(t5_model_name).eval().requires_grad_(False).to(torch_dtype))
|
||||
|
||||
def forward(self, texts: list[str], device: torch.device | str) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
encoded = self.tokenizer(texts, truncation=True, max_length=self.max_length,
|
||||
padding="max_length", return_tensors="pt")
|
||||
input_ids = encoded["input_ids"].to(device)
|
||||
attention_mask = encoded["attention_mask"].to(device).to(torch.bool)
|
||||
# Mirror official's `autocast(fp16)` wrap on T5 forward.
|
||||
with torch.no_grad(), torch.autocast(device_type="cuda", dtype=self._t5_dtype):
|
||||
embeddings = self.model(input_ids=input_ids,
|
||||
attention_mask=attention_mask)["last_hidden_state"]
|
||||
embeddings = self.proj_out(embeddings) * attention_mask.unsqueeze(-1).float()
|
||||
return embeddings, attention_mask
|
||||
|
||||
|
||||
class NumberConditioner(_Conditioner):
|
||||
"""Float-valued conditioner with min/max clamping + NumberEmbedder."""
|
||||
|
||||
def __init__(self, output_dim: int, min_val: float = 0, max_val: float = 1) -> None:
|
||||
super().__init__(output_dim, output_dim)
|
||||
self.min_val = min_val
|
||||
self.max_val = max_val
|
||||
self.embedder = NumberEmbedder(features=output_dim)
|
||||
|
||||
def forward(self, floats: list[float], device: torch.device | str) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
floats = [float(x) for x in floats]
|
||||
floats_t = torch.tensor(floats, device=device).clamp(self.min_val, self.max_val)
|
||||
normalized = (floats_t - self.min_val) / (self.max_val - self.min_val)
|
||||
emb_dtype = next(self.embedder.parameters()).dtype
|
||||
normalized = normalized.to(emb_dtype)
|
||||
float_embeds = self.embedder(normalized).unsqueeze(1)
|
||||
return float_embeds, torch.ones(float_embeds.shape[0], 1, device=device)
|
||||
|
||||
|
||||
class StableAudioMultiConditioner(nn.Module):
|
||||
"""SA-Open-1.0 conditioner: T5 prompt + duration NumberConditioners.
|
||||
|
||||
All hardcoded constants (cond_dim, sub-conditioner ids, T5 model
|
||||
name + max_length, NumberConditioner ranges) live on
|
||||
`StableAudioConditionerConfig` — see
|
||||
`fastvideo/configs/models/encoders/stable_audio_conditioner.py`.
|
||||
"""
|
||||
|
||||
def __init__(self, config: StableAudioConditionerConfig | None = None) -> None:
|
||||
super().__init__()
|
||||
self.config = config or StableAudioConditionerConfig()
|
||||
arch = self.config.arch_config
|
||||
# Build sub-conditioners from the `configs` list (mirrors
|
||||
# upstream's `MultiConditioner` factory).
|
||||
sub: dict[str, nn.Module] = {}
|
||||
for spec in arch.configs:
|
||||
sid = spec["id"]
|
||||
stype = spec["type"]
|
||||
scfg = spec["config"]
|
||||
if stype == "t5":
|
||||
sub[sid] = T5Conditioner(output_dim=arch.cond_dim,
|
||||
t5_model_name=scfg["t5_model_name"],
|
||||
max_length=scfg["max_length"],
|
||||
dtype=arch.t5_dtype)
|
||||
elif stype == "number":
|
||||
sub[sid] = NumberConditioner(output_dim=arch.cond_dim,
|
||||
min_val=scfg["min_val"], max_val=scfg["max_val"])
|
||||
else:
|
||||
raise ValueError(f"Unknown sub-conditioner type {stype!r} for id {sid!r}.")
|
||||
self.conditioners = nn.ModuleDict(sub)
|
||||
self.cross_attention_cond_ids = tuple(arch.cross_attention_cond_ids)
|
||||
self.global_cond_ids = tuple(arch.global_cond_ids)
|
||||
|
||||
def forward(self, batch_metadata: list[dict],
|
||||
device: torch.device | str) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
|
||||
out: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
|
||||
for key, conditioner in self.conditioners.items():
|
||||
inputs = [x[key] for x in batch_metadata]
|
||||
out[key] = conditioner(inputs, device)
|
||||
return out
|
||||
|
||||
def get_conditioning_inputs(
|
||||
self, cond: dict[str, tuple[torch.Tensor, torch.Tensor]]
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Pack conditioner outputs into the (cross_attn_cond,
|
||||
cross_attn_mask, global_embed) triple the DiT consumes. Order
|
||||
is driven by `cross_attention_cond_ids` / `global_cond_ids`
|
||||
from the config — SA-1.0 uses three sub-conditioners
|
||||
(prompt + seconds_start + seconds_total); SA-small uses two
|
||||
(prompt + seconds_total).
|
||||
"""
|
||||
x_embs = [cond[i][0] for i in self.cross_attention_cond_ids]
|
||||
x_masks = [cond[i][1] for i in self.cross_attention_cond_ids]
|
||||
cross_attn_cond = torch.cat(x_embs, dim=1)
|
||||
cross_attn_mask = torch.cat(x_masks, dim=1)
|
||||
global_embed = torch.cat([cond[i][0][:, 0] for i in self.global_cond_ids], dim=-1)
|
||||
return cross_attn_cond, cross_attn_mask, global_embed
|
||||
|
||||
@classmethod
|
||||
def from_official_state_dict(cls, state_dict: dict[str, torch.Tensor],
|
||||
prefix: str = "conditioner.") -> "StableAudioMultiConditioner":
|
||||
"""Load NumberConditioner weights from a raw `stable_audio_tools`
|
||||
monolithic state dict. Kept for tests / older checkpoints;
|
||||
production loads go through the standard `ConditionerLoader`
|
||||
against the converted Diffusers repo.
|
||||
"""
|
||||
mc = cls()
|
||||
own_state = mc.state_dict()
|
||||
loaded: dict[str, torch.Tensor] = {}
|
||||
for k, v in state_dict.items():
|
||||
if not k.startswith(prefix):
|
||||
continue
|
||||
stripped = k[len(prefix):]
|
||||
if stripped in own_state:
|
||||
loaded[stripped] = v
|
||||
# T5 keys are intentionally absent from the checkpoint.
|
||||
missing = [k for k in own_state.keys() if k not in loaded
|
||||
and not k.startswith("conditioners.prompt.")]
|
||||
unexpected = [k for k in loaded.keys() if k not in own_state]
|
||||
if missing or unexpected:
|
||||
raise RuntimeError(
|
||||
f"StableAudioMultiConditioner load mismatch — missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
mc.load_state_dict(loaded, strict=False)
|
||||
return mc
|
||||
|
||||
|
||||
EntryClass = StableAudioMultiConditioner
|
||||
@@ -0,0 +1,224 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Inspired by SGLang's layerwise offload implementation:
|
||||
# https://github.com/sgl-project/sglang/pull/15511
|
||||
#
|
||||
# This implementation provides a lightweight layerwise CPU offload manager
|
||||
# with async H2D prefetch using a dedicated CUDA stream, following SGLang's design.
|
||||
|
||||
import re
|
||||
from contextlib import contextmanager
|
||||
from typing import Dict, Set, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class LayerwiseOffloadManager:
|
||||
"""A lightweight layerwise CPU offload manager.
|
||||
|
||||
Offloads per-layer parameters/buffers from GPU to CPU, and supports async H2D
|
||||
prefetch using a dedicated CUDA stream.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: torch.nn.Module,
|
||||
*,
|
||||
module_list_attr: str,
|
||||
num_layers: int,
|
||||
enabled: bool,
|
||||
pin_cpu_memory: bool = True,
|
||||
auto_initialize: bool = False,
|
||||
) -> None:
|
||||
self.model = model
|
||||
self.module_list_attr = module_list_attr
|
||||
self.num_layers = int(num_layers)
|
||||
self.pin_cpu_memory = bool(pin_cpu_memory)
|
||||
|
||||
self.enabled = bool(enabled and torch.cuda.is_available())
|
||||
self.device = (
|
||||
torch.device("cuda", torch.cuda.current_device()) if self.enabled else None
|
||||
)
|
||||
self.copy_stream = torch.cuda.Stream() if self.enabled else None
|
||||
|
||||
self._layer_name_re = re.compile(
|
||||
rf"(^|\.){re.escape(module_list_attr)}\.(\d+)(\.|$)"
|
||||
)
|
||||
|
||||
self._cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
|
||||
self._cpu_dtypes: Dict[int, Dict[str, torch.dtype]] = {}
|
||||
|
||||
self._gpu_layers: Dict[int, Set[str]] = {}
|
||||
|
||||
self._named_parameters: Dict[str, torch.nn.Parameter] = {}
|
||||
self._named_buffers: Dict[str, torch.Tensor] = {}
|
||||
|
||||
self._meta: Dict[str, Tuple[int, torch.dtype]] = {}
|
||||
|
||||
if auto_initialize:
|
||||
self.initialize()
|
||||
|
||||
def _match_layer_idx(self, name: str) -> Optional[int]:
|
||||
m = self._layer_name_re.search(name)
|
||||
if not m:
|
||||
return None
|
||||
try:
|
||||
return int(m.group(2))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _record_meta(self, name: str, t: torch.Tensor) -> None:
|
||||
if name not in self._meta:
|
||||
self._meta[name] = (int(t.ndim), t.dtype)
|
||||
|
||||
def _make_placeholder(self, name: str) -> torch.Tensor:
|
||||
"""Rank-preserving empty placeholder on GPU."""
|
||||
assert self.device is not None
|
||||
ndim, dtype = self._meta[name]
|
||||
shape = (0,) if ndim <= 0 else (0,) * ndim
|
||||
return torch.empty(shape, device=self.device, dtype=dtype)
|
||||
|
||||
def _get_target(self, name: str) -> torch.Tensor:
|
||||
if name in self._named_parameters:
|
||||
return self._named_parameters[name]
|
||||
return self._named_buffers[name]
|
||||
|
||||
def _offload_tensor(self, name: str, tensor: torch.Tensor, layer_idx: int) -> None:
|
||||
if layer_idx not in self._cpu_weights:
|
||||
self._cpu_weights[layer_idx] = {}
|
||||
self._cpu_dtypes[layer_idx] = {}
|
||||
|
||||
self._record_meta(name, tensor)
|
||||
|
||||
cpu_weight = tensor.detach().to("cpu")
|
||||
if self.pin_cpu_memory:
|
||||
cpu_weight = cpu_weight.pin_memory()
|
||||
|
||||
self._cpu_weights[layer_idx][name] = cpu_weight
|
||||
self._cpu_dtypes[layer_idx][name] = tensor.dtype
|
||||
|
||||
if self.device is not None:
|
||||
tensor.data = self._make_placeholder(name)
|
||||
|
||||
@torch.compiler.disable
|
||||
def initialize(self) -> None:
|
||||
"""Offload all matched layer tensors to CPU and prefetch layer 0 (sync)."""
|
||||
if not self.enabled:
|
||||
return
|
||||
|
||||
self._named_parameters = dict(self.model.named_parameters())
|
||||
self._named_buffers = dict(self.model.named_buffers())
|
||||
|
||||
for name, param in self._named_parameters.items():
|
||||
layer_idx = self._match_layer_idx(name)
|
||||
if layer_idx is None or layer_idx >= self.num_layers:
|
||||
continue
|
||||
self._offload_tensor(name, param, layer_idx)
|
||||
|
||||
for name, buf in self._named_buffers.items():
|
||||
layer_idx = self._match_layer_idx(name)
|
||||
if layer_idx is None or layer_idx >= self.num_layers:
|
||||
continue
|
||||
self._offload_tensor(name, buf, layer_idx)
|
||||
|
||||
self.prefetch_layer(0, non_blocking=False)
|
||||
if self.copy_stream is not None:
|
||||
torch.cuda.current_stream().wait_stream(self.copy_stream)
|
||||
|
||||
@torch.compiler.disable
|
||||
def prefetch_layer(self, layer_idx: int, non_blocking: bool = True) -> None:
|
||||
"""Prefetch a layer's tensors from CPU to GPU (async on copy_stream)."""
|
||||
if not self.enabled or self.device is None or self.copy_stream is None:
|
||||
return
|
||||
if layer_idx < 0 or layer_idx >= self.num_layers:
|
||||
return
|
||||
if layer_idx in self._gpu_layers:
|
||||
return
|
||||
if layer_idx not in self._cpu_weights:
|
||||
return
|
||||
|
||||
self.copy_stream.wait_stream(torch.cuda.current_stream())
|
||||
|
||||
param_names: Set[str] = set()
|
||||
with torch.cuda.stream(self.copy_stream):
|
||||
for name, cpu_weight in self._cpu_weights[layer_idx].items():
|
||||
target = self._get_target(name)
|
||||
|
||||
gpu_weight = torch.empty(
|
||||
cpu_weight.shape,
|
||||
dtype=self._cpu_dtypes[layer_idx][name],
|
||||
device=self.device,
|
||||
)
|
||||
gpu_weight.copy_(cpu_weight, non_blocking=non_blocking)
|
||||
|
||||
target.data = gpu_weight
|
||||
param_names.add(name)
|
||||
|
||||
self._gpu_layers[layer_idx] = param_names
|
||||
|
||||
@contextmanager
|
||||
def layer_scope(
|
||||
self,
|
||||
*,
|
||||
prefetch_layer_idx: Optional[int],
|
||||
release_layer_idx: Optional[int],
|
||||
non_blocking: bool = True,
|
||||
):
|
||||
if self.enabled and release_layer_idx is not None:
|
||||
cur = release_layer_idx
|
||||
if (
|
||||
cur not in self._gpu_layers
|
||||
and cur in self._cpu_weights
|
||||
and self.device is not None
|
||||
and self.copy_stream is not None
|
||||
):
|
||||
self.prefetch_layer(cur, non_blocking=False)
|
||||
torch.cuda.current_stream().wait_stream(self.copy_stream)
|
||||
|
||||
if self.enabled and prefetch_layer_idx is not None:
|
||||
self.prefetch_layer(prefetch_layer_idx, non_blocking=non_blocking)
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if self.enabled and self.copy_stream is not None:
|
||||
torch.cuda.current_stream().wait_stream(self.copy_stream)
|
||||
if self.enabled and release_layer_idx is not None:
|
||||
self.release_layer(release_layer_idx)
|
||||
|
||||
|
||||
@torch.compiler.disable
|
||||
def release_layer(self, layer_idx: int) -> None:
|
||||
"""Release a layer's tensors back to placeholders (free VRAM)."""
|
||||
if not self.enabled or self.device is None:
|
||||
return
|
||||
|
||||
if layer_idx < 0:
|
||||
return
|
||||
|
||||
param_names = self._gpu_layers.pop(layer_idx, None)
|
||||
if not param_names:
|
||||
return
|
||||
|
||||
for name in param_names:
|
||||
target = self._get_target(name)
|
||||
# Ensure meta exists even if something unexpected happened
|
||||
self._record_meta(name, target)
|
||||
target.data = self._make_placeholder(name)
|
||||
|
||||
@torch.compiler.disable
|
||||
def release_all(self) -> None:
|
||||
"""Release all currently-resident layers back to placeholders."""
|
||||
if not self.enabled or self.device is None:
|
||||
return
|
||||
|
||||
if self.copy_stream is not None:
|
||||
torch.cuda.current_stream().wait_stream(self.copy_stream)
|
||||
|
||||
for layer_idx in list(self._gpu_layers.keys()):
|
||||
param_names = self._gpu_layers.pop(layer_idx, None)
|
||||
if not param_names:
|
||||
continue
|
||||
for name in param_names:
|
||||
target = self._get_target(name)
|
||||
self._record_meta(name, target)
|
||||
target.data = self._make_placeholder(name)
|
||||
@@ -95,10 +95,6 @@ class ComponentLoader(ABC):
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
"upsampler": (UpsamplerLoader, "diffusers"),
|
||||
"upsampler_2": (UpsamplerLoader, "diffusers"),
|
||||
# Stable Audio's `StableAudioMultiConditioner` bundles T5 +
|
||||
# NumberConditioners; not a pure text encoder, so it gets
|
||||
# its own loader.
|
||||
"conditioner": (ConditionerLoader, "fastvideo"),
|
||||
}
|
||||
|
||||
if module_type in module_loaders:
|
||||
@@ -1002,61 +998,6 @@ class SchedulerLoader(ComponentLoader):
|
||||
return scheduler
|
||||
|
||||
|
||||
class ConditionerLoader(ComponentLoader):
|
||||
"""Loader for multi-conditioner components (e.g. Stable Audio's
|
||||
`StableAudioMultiConditioner`, which bundles T5 + NumberConditioners
|
||||
and is neither a pure text encoder nor a Diffusers-shaped module).
|
||||
Reads `<subfolder>/config.json` to resolve the class via
|
||||
`ModelRegistry`, instantiates with no args (the class pulls its own
|
||||
defaults from its FastVideo config), then loads
|
||||
`diffusion_pytorch_model.safetensors` non-strictly so externally
|
||||
fetched sub-encoders (T5) don't trip the missing-key check.
|
||||
"""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
config = get_diffusers_config(model=model_path)
|
||||
class_name = config.pop("_class_name", None)
|
||||
config.pop("_name_or_path", None)
|
||||
if class_name is None:
|
||||
raise ValueError(
|
||||
f"Conditioner config at {model_path} is missing the "
|
||||
f"`_class_name` attribute required to resolve a model class.")
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
target_device = get_local_torch_device()
|
||||
precision = getattr(fastvideo_args.pipeline_config, "precision", "fp16")
|
||||
target_dtype = PRECISION_TO_TYPE.get(precision, torch.float16)
|
||||
|
||||
# Without this merge the model falls back to its dataclass
|
||||
# defaults (e.g. SA-1.0's 3-conditioner spec — wrong for SA-small).
|
||||
from dataclasses import fields as _fields
|
||||
from fastvideo.configs.models.encoders import (
|
||||
StableAudioConditionerConfig, )
|
||||
if model_cls.__name__ == "StableAudioMultiConditioner":
|
||||
cond_config = StableAudioConditionerConfig()
|
||||
# `update_model_arch` is strict (raises on unknown keys); the
|
||||
# converter writes a few non-arch keys (`_class_name`,
|
||||
# `_diffusers_version`, `_name_or_path`) that must be filtered
|
||||
# out first.
|
||||
valid = {f.name for f in _fields(cond_config.arch_config)}
|
||||
cond_config.update_model_arch({k: v for k, v in config.items() if k in valid})
|
||||
with set_default_torch_dtype(target_dtype):
|
||||
model = model_cls(cond_config)
|
||||
else:
|
||||
with set_default_torch_dtype(target_dtype):
|
||||
model = model_cls()
|
||||
|
||||
weights = os.path.join(str(model_path), "diffusion_pytorch_model.safetensors")
|
||||
if not os.path.isfile(weights):
|
||||
raise FileNotFoundError(
|
||||
f"Conditioner weights not found: {weights}")
|
||||
state = safetensors_load_file(weights)
|
||||
# Non-strict: T5 weights live outside this checkpoint (fetched in
|
||||
# the conditioner's `__init__` from the standard HF repo).
|
||||
model.load_state_dict(state, strict=False)
|
||||
return model.to(device=target_device, dtype=target_dtype).eval()
|
||||
|
||||
|
||||
class UpsamplerLoader(ComponentLoader):
|
||||
"""Loader for upsamplers."""
|
||||
|
||||
|
||||
@@ -86,9 +86,6 @@ _VAE_MODELS = {
|
||||
("vaes", "gen3c_tokenizer_vae", "AutoencoderKLGen3CTokenizer"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
|
||||
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
|
||||
# `stable-audio-open-1.0/vae/config.json` ships `_class_name="AutoencoderOobleck"`
|
||||
# (Diffusers' name); FastVideo's class is `OobleckVAE`.
|
||||
"AutoencoderOobleck": ("vaes", "oobleck", "OobleckVAE"),
|
||||
}
|
||||
|
||||
_AUDIO_MODELS = {
|
||||
|
||||
@@ -1,376 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 "Oobleck" VAE.
|
||||
|
||||
5-stage Conv1d autoencoder with Snake activations + diagonal-Gaussian
|
||||
bottleneck. Loads `stabilityai/stable-audio-open-1.0/vae/` weights
|
||||
directly via `OobleckVAE.from_pretrained(...)`.
|
||||
|
||||
vae = OobleckVAE.from_pretrained("stabilityai/stable-audio-open-1.0", subfolder="vae")
|
||||
waveform = vae.decode(latent) # (B, audio_channels, samples)
|
||||
latent = vae.encode(waveform).sample() # or .mode()
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Snake1d(nn.Module):
|
||||
"""A 1D Snake activation with learnable per-channel alpha/beta."""
|
||||
|
||||
def __init__(self, hidden_dim: int, logscale: bool = True):
|
||||
super().__init__()
|
||||
self.alpha = nn.Parameter(torch.zeros(1, hidden_dim, 1))
|
||||
self.beta = nn.Parameter(torch.zeros(1, hidden_dim, 1))
|
||||
self.alpha.requires_grad = True
|
||||
self.beta.requires_grad = True
|
||||
self.logscale = logscale
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
shape = x.shape
|
||||
alpha = self.alpha if not self.logscale else torch.exp(self.alpha)
|
||||
beta = self.beta if not self.logscale else torch.exp(self.beta)
|
||||
x = x.reshape(shape[0], shape[1], -1)
|
||||
x = x + (beta + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
|
||||
return x.reshape(shape)
|
||||
|
||||
|
||||
class OobleckResidualUnit(nn.Module):
|
||||
def __init__(self, dimension: int = 16, dilation: int = 1):
|
||||
super().__init__()
|
||||
pad = ((7 - 1) * dilation) // 2
|
||||
self.snake1 = Snake1d(dimension)
|
||||
self.conv1 = weight_norm(nn.Conv1d(
|
||||
dimension, dimension, kernel_size=7, dilation=dilation, padding=pad,
|
||||
))
|
||||
self.snake2 = Snake1d(dimension)
|
||||
self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = self.conv1(self.snake1(x))
|
||||
out = self.conv2(self.snake2(out))
|
||||
pad = (x.shape[-1] - out.shape[-1]) // 2
|
||||
if pad > 0:
|
||||
x = x[..., pad:-pad]
|
||||
return x + out
|
||||
|
||||
|
||||
class OobleckEncoderBlock(nn.Module):
|
||||
def __init__(self, input_dim: int, output_dim: int, stride: int = 1):
|
||||
super().__init__()
|
||||
self.res_unit1 = OobleckResidualUnit(input_dim, dilation=1)
|
||||
self.res_unit2 = OobleckResidualUnit(input_dim, dilation=3)
|
||||
self.res_unit3 = OobleckResidualUnit(input_dim, dilation=9)
|
||||
self.snake1 = Snake1d(input_dim)
|
||||
self.conv1 = weight_norm(nn.Conv1d(
|
||||
input_dim, output_dim,
|
||||
kernel_size=2 * stride, stride=stride,
|
||||
padding=math.ceil(stride / 2),
|
||||
))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.res_unit1(x)
|
||||
x = self.res_unit2(x)
|
||||
x = self.snake1(self.res_unit3(x))
|
||||
return self.conv1(x)
|
||||
|
||||
|
||||
class OobleckDecoderBlock(nn.Module):
|
||||
def __init__(self, input_dim: int, output_dim: int, stride: int = 1):
|
||||
super().__init__()
|
||||
self.snake1 = Snake1d(input_dim)
|
||||
self.conv_t1 = weight_norm(nn.ConvTranspose1d(
|
||||
input_dim, output_dim,
|
||||
kernel_size=2 * stride, stride=stride,
|
||||
padding=math.ceil(stride / 2),
|
||||
))
|
||||
self.res_unit1 = OobleckResidualUnit(output_dim, dilation=1)
|
||||
self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3)
|
||||
self.res_unit3 = OobleckResidualUnit(output_dim, dilation=9)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.snake1(x)
|
||||
x = self.conv_t1(x)
|
||||
x = self.res_unit1(x)
|
||||
x = self.res_unit2(x)
|
||||
return self.res_unit3(x)
|
||||
|
||||
|
||||
class OobleckDiagonalGaussianDistribution:
|
||||
"""Diagonal-Gaussian VAE posterior with `softplus(scale) + 1e-4` std."""
|
||||
|
||||
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
|
||||
self.parameters = parameters
|
||||
self.mean, self.scale = parameters.chunk(2, dim=1)
|
||||
self.std = nn.functional.softplus(self.scale) + 1e-4
|
||||
self.var = self.std * self.std
|
||||
self.logvar = torch.log(self.var)
|
||||
self.deterministic = deterministic
|
||||
|
||||
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
noise = torch.randn(
|
||||
self.mean.shape, generator=generator,
|
||||
device=self.parameters.device, dtype=self.parameters.dtype,
|
||||
)
|
||||
return self.mean + self.std * noise
|
||||
|
||||
def mode(self) -> torch.Tensor:
|
||||
return self.mean
|
||||
|
||||
|
||||
@dataclass
|
||||
class OobleckDecoderOutput:
|
||||
sample: torch.Tensor
|
||||
|
||||
|
||||
class OobleckEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
encoder_hidden_size: int,
|
||||
audio_channels: int,
|
||||
downsampling_ratios: list[int],
|
||||
channel_multiples: list[int],
|
||||
):
|
||||
super().__init__()
|
||||
strides = downsampling_ratios
|
||||
channel_multiples = [1] + list(channel_multiples)
|
||||
self.conv1 = weight_norm(nn.Conv1d(
|
||||
audio_channels, encoder_hidden_size, kernel_size=7, padding=3,
|
||||
))
|
||||
self.block = nn.ModuleList([
|
||||
OobleckEncoderBlock(
|
||||
input_dim=encoder_hidden_size * channel_multiples[i],
|
||||
output_dim=encoder_hidden_size * channel_multiples[i + 1],
|
||||
stride=s,
|
||||
)
|
||||
for i, s in enumerate(strides)
|
||||
])
|
||||
d_model = encoder_hidden_size * channel_multiples[-1]
|
||||
self.snake1 = Snake1d(d_model)
|
||||
self.conv2 = weight_norm(nn.Conv1d(
|
||||
d_model, encoder_hidden_size, kernel_size=3, padding=1,
|
||||
))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.conv1(x)
|
||||
for m in self.block:
|
||||
x = m(x)
|
||||
x = self.snake1(x)
|
||||
return self.conv2(x)
|
||||
|
||||
|
||||
class OobleckDecoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
input_channels: int,
|
||||
audio_channels: int,
|
||||
upsampling_ratios: list[int],
|
||||
channel_multiples: list[int],
|
||||
):
|
||||
super().__init__()
|
||||
strides = upsampling_ratios
|
||||
channel_multiples = [1] + list(channel_multiples)
|
||||
self.conv1 = weight_norm(nn.Conv1d(
|
||||
input_channels, channels * channel_multiples[-1],
|
||||
kernel_size=7, padding=3,
|
||||
))
|
||||
self.block = nn.ModuleList([
|
||||
OobleckDecoderBlock(
|
||||
input_dim=channels * channel_multiples[len(strides) - i],
|
||||
output_dim=channels * channel_multiples[len(strides) - i - 1],
|
||||
stride=s,
|
||||
)
|
||||
for i, s in enumerate(strides)
|
||||
])
|
||||
self.snake1 = Snake1d(channels)
|
||||
self.conv2 = weight_norm(nn.Conv1d(
|
||||
channels, audio_channels, kernel_size=7, padding=3, bias=False,
|
||||
))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.conv1(x)
|
||||
for layer in self.block:
|
||||
x = layer(x)
|
||||
x = self.snake1(x)
|
||||
return self.conv2(x)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Top-level VAE
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class OobleckVAE(nn.Module):
|
||||
"""Stable Audio Open 1.0 VAE.
|
||||
|
||||
Constructed either from an `OobleckVAEConfig` (the standard
|
||||
`VAELoader` path) or from explicit kwargs (back-compat for tests
|
||||
and `from_pretrained` callers).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config=None, # type: OobleckVAEConfig | None
|
||||
*,
|
||||
encoder_hidden_size: int = 128,
|
||||
downsampling_ratios: list[int] | None = None,
|
||||
channel_multiples: list[int] | None = None,
|
||||
decoder_channels: int = 128,
|
||||
decoder_input_channels: int = 64,
|
||||
audio_channels: int = 2,
|
||||
sampling_rate: int = 44100,
|
||||
):
|
||||
super().__init__()
|
||||
if config is not None:
|
||||
arch = config.arch_config
|
||||
encoder_hidden_size = arch.encoder_hidden_size
|
||||
downsampling_ratios = list(arch.downsampling_ratios)
|
||||
channel_multiples = list(arch.channel_multiples)
|
||||
decoder_channels = arch.decoder_channels
|
||||
decoder_input_channels = arch.decoder_input_channels
|
||||
audio_channels = arch.audio_channels
|
||||
sampling_rate = arch.sampling_rate
|
||||
if downsampling_ratios is None:
|
||||
downsampling_ratios = [2, 4, 4, 8, 8]
|
||||
if channel_multiples is None:
|
||||
channel_multiples = [1, 2, 4, 8, 16]
|
||||
self.encoder_hidden_size = encoder_hidden_size
|
||||
self.downsampling_ratios = downsampling_ratios
|
||||
self.decoder_channels = decoder_channels
|
||||
self.upsampling_ratios = list(reversed(downsampling_ratios))
|
||||
self.hop_length = int(np.prod(downsampling_ratios))
|
||||
self.sampling_rate = sampling_rate
|
||||
self.audio_channels = audio_channels
|
||||
self.decoder_input_channels = decoder_input_channels
|
||||
|
||||
self.encoder = OobleckEncoder(
|
||||
encoder_hidden_size=encoder_hidden_size,
|
||||
audio_channels=audio_channels,
|
||||
downsampling_ratios=downsampling_ratios,
|
||||
channel_multiples=channel_multiples,
|
||||
)
|
||||
self.decoder = OobleckDecoder(
|
||||
channels=decoder_channels,
|
||||
input_channels=decoder_input_channels,
|
||||
audio_channels=audio_channels,
|
||||
upsampling_ratios=self.upsampling_ratios,
|
||||
channel_multiples=channel_multiples,
|
||||
)
|
||||
|
||||
def encode(
|
||||
self, x: torch.Tensor,
|
||||
) -> OobleckDiagonalGaussianDistribution:
|
||||
return OobleckDiagonalGaussianDistribution(self.encoder(x))
|
||||
|
||||
def decode(self, z: torch.Tensor) -> OobleckDecoderOutput:
|
||||
return OobleckDecoderOutput(sample=self.decoder(z))
|
||||
|
||||
def forward(
|
||||
self, sample: torch.Tensor, sample_posterior: bool = False,
|
||||
) -> OobleckDecoderOutput:
|
||||
posterior = self.encode(sample)
|
||||
z = posterior.sample() if sample_posterior else posterior.mode()
|
||||
return self.decode(z)
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Loader
|
||||
# -------------------------------------------------------------------
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls,
|
||||
model_path: str,
|
||||
*,
|
||||
subfolder: str | None = None,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
) -> "OobleckVAE":
|
||||
"""Instantiate and load weights from a Stable Audio VAE dir.
|
||||
|
||||
`model_path` may be:
|
||||
* a HF repo id (e.g. `stabilityai/stable-audio-open-1.0`),
|
||||
* a local directory containing `config.json` + safetensors,
|
||||
* a local directory whose `subfolder="vae"` holds those files.
|
||||
|
||||
For gated repos, the HF token is read from `HF_TOKEN` /
|
||||
`HUGGINGFACE_HUB_TOKEN` / `HF_API_KEY` (see `resolve_hf_token`).
|
||||
"""
|
||||
import inspect
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.utils import resolve_hf_token
|
||||
|
||||
# Resolve to a local directory.
|
||||
if os.path.isdir(model_path):
|
||||
root = model_path
|
||||
else:
|
||||
from huggingface_hub import snapshot_download
|
||||
allow = ["vae/*"] if subfolder else ["*"]
|
||||
root = snapshot_download(
|
||||
repo_id=model_path, token=resolve_hf_token(), allow_patterns=allow,
|
||||
)
|
||||
if subfolder:
|
||||
root = os.path.join(root, subfolder)
|
||||
if not os.path.isdir(root):
|
||||
raise FileNotFoundError(f"Not a directory: {root}")
|
||||
|
||||
cfg_path = os.path.join(root, "config.json")
|
||||
if not os.path.isfile(cfg_path):
|
||||
raise FileNotFoundError(
|
||||
f"Expected config.json at {cfg_path}. If using a HF repo, "
|
||||
f"pass subfolder='vae'."
|
||||
)
|
||||
with open(cfg_path) as f:
|
||||
cfg = json.load(f)
|
||||
cfg_fields = {k: v for k, v in cfg.items() if not k.startswith("_")}
|
||||
# Diffusers configs commonly carry extra fields (`scaling_factor`,
|
||||
# `_diffusers_version`, ...) the bare `OobleckVAE` ctor doesn't accept.
|
||||
init_params = inspect.signature(cls.__init__).parameters
|
||||
cfg_fields = {k: v for k, v in cfg_fields.items() if k in init_params}
|
||||
|
||||
model = cls(**cfg_fields)
|
||||
|
||||
weights_path = os.path.join(root, "diffusion_pytorch_model.safetensors")
|
||||
if not os.path.isfile(weights_path):
|
||||
# Allow `model.safetensors` as a fallback.
|
||||
alt = os.path.join(root, "model.safetensors")
|
||||
if os.path.isfile(alt):
|
||||
weights_path = alt
|
||||
else:
|
||||
raise FileNotFoundError(
|
||||
f"No safetensors weights under {root}. Expected "
|
||||
f"diffusion_pytorch_model.safetensors."
|
||||
)
|
||||
state = load_file(weights_path)
|
||||
missing, unexpected = model.load_state_dict(state, strict=False)
|
||||
if missing:
|
||||
raise RuntimeError(
|
||||
f"OobleckVAE missing {len(missing)} keys from {weights_path}: "
|
||||
f"{missing[:5]}"
|
||||
)
|
||||
if unexpected:
|
||||
# Non-critical: some checkpoints embed the VAE inside a larger
|
||||
# container (e.g. `pretransform.model.*`). Log the count so
|
||||
# genuine loader regressions don't go unnoticed.
|
||||
logger.debug(
|
||||
"OobleckVAE: ignored %d unexpected keys from %s "
|
||||
"(first 3: %s)", len(unexpected), weights_path, unexpected[:3],
|
||||
)
|
||||
|
||||
if torch_dtype is not None:
|
||||
model = model.to(dtype=torch_dtype)
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
EntryClass = OobleckVAE
|
||||
@@ -1,122 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Lazy-loading pipeline wrapper around `OobleckVAE`.
|
||||
|
||||
Two reasons this exists rather than using `OobleckVAE` directly:
|
||||
|
||||
1. The underlying VAE is fetched on first `encode`/`decode` call, not
|
||||
at construction — lets pipelines build the module tree on CPU
|
||||
before knowing the target device.
|
||||
2. The lazy VAE's params are hidden from `named_parameters()` so the
|
||||
FastVideo pipeline-component loader doesn't try to match Oobleck's
|
||||
safetensors against the host pipeline's converted-repo state dict.
|
||||
|
||||
For standalone use prefer `OobleckVAE.from_pretrained(...)` directly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.configs.models.vaes import OobleckVAEConfig
|
||||
|
||||
|
||||
class SAAudioVAEModel(nn.Module):
|
||||
"""Pipeline-glue lazy loader around `OobleckVAE`."""
|
||||
|
||||
def __init__(self, config: OobleckVAEConfig) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
arch = config.arch_config
|
||||
self.pretrained_path: str = config.pretrained_path
|
||||
self.pretrained_subfolder: str | None = config.pretrained_subfolder
|
||||
self.pretrained_dtype: str = config.pretrained_dtype
|
||||
self.sampling_rate: int = arch.sampling_rate
|
||||
self.audio_channels: int = arch.audio_channels
|
||||
self.decoder_input_channels: int = arch.decoder_input_channels
|
||||
self._oobleck_vae = None
|
||||
|
||||
def named_parameters(self, prefix: str = "", recurse: bool = True):
|
||||
# Hide the lazy-loaded VAE — its weights are fetched separately
|
||||
# and shouldn't appear in the host pipeline's loader sweep.
|
||||
for name, param in super().named_parameters(prefix=prefix, recurse=recurse):
|
||||
if name.startswith("_oobleck_vae.") or name == "_oobleck_vae":
|
||||
continue
|
||||
yield name, param
|
||||
|
||||
def _build(self, device: torch.device | None = None):
|
||||
from fastvideo.models.vaes.oobleck import OobleckVAE
|
||||
|
||||
path = self.pretrained_path
|
||||
if not path:
|
||||
raise ValueError(
|
||||
"OobleckVAEConfig.pretrained_path must be set; expected "
|
||||
"`stabilityai/stable-audio-open-1.0` or a local path."
|
||||
)
|
||||
dtype = getattr(torch, self.pretrained_dtype, torch.float32)
|
||||
# If the caller already pointed us at the VAE dir directly, drop
|
||||
# the subfolder. Otherwise pass through (default "vae").
|
||||
subfolder: str | None = self.pretrained_subfolder
|
||||
if subfolder and os.path.isdir(path) and os.path.isfile(os.path.join(path, "config.json")):
|
||||
subfolder = None
|
||||
model = OobleckVAE.from_pretrained(path, subfolder=subfolder, torch_dtype=dtype)
|
||||
if device is not None:
|
||||
model = model.to(device=device)
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
@property
|
||||
def oobleck_vae(self):
|
||||
if self._oobleck_vae is None:
|
||||
self._oobleck_vae = self._build()
|
||||
return self._oobleck_vae
|
||||
|
||||
# Back-compat alias: callers that imported this earlier referred to
|
||||
# the underlying VAE as `sa_audio_vae_model`. Both names point at the
|
||||
# same object.
|
||||
@property
|
||||
def sa_audio_vae_model(self):
|
||||
return self.oobleck_vae
|
||||
|
||||
@property
|
||||
def hop_length(self) -> int:
|
||||
return int(self.oobleck_vae.hop_length)
|
||||
|
||||
def _move_to_input_device(self, model, ref: torch.Tensor):
|
||||
if ref is None:
|
||||
return model
|
||||
first_param = next(model.parameters(), None)
|
||||
if first_param is not None and first_param.device != ref.device:
|
||||
model = model.to(device=ref.device)
|
||||
self._oobleck_vae = model
|
||||
return model
|
||||
|
||||
def decode(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Decode an audio latent (`[B, C_latent, L]`) -> waveform
|
||||
(`[B, audio_channels, samples]`).
|
||||
"""
|
||||
model = self.oobleck_vae
|
||||
model = self._move_to_input_device(model, latent)
|
||||
with torch.no_grad():
|
||||
out = model.decode(latent.to(next(model.parameters()).dtype))
|
||||
if hasattr(out, "sample"):
|
||||
return out.sample
|
||||
return out
|
||||
|
||||
def encode(self, waveform: torch.Tensor, sample_posterior: bool = False) -> torch.Tensor:
|
||||
"""Encode `[B, C_audio, samples]` -> latent `[B, C_latent, L]`.
|
||||
|
||||
`sample_posterior=False` (default): deterministic mean.
|
||||
`sample_posterior=True`: stochastic sample (`mean + softplus(scale) * randn`).
|
||||
"""
|
||||
model = self.oobleck_vae
|
||||
model = self._move_to_input_device(model, waveform)
|
||||
with torch.no_grad():
|
||||
out = model.encode(waveform.to(next(model.parameters()).dtype))
|
||||
if hasattr(out, "latent_dist"):
|
||||
out = out.latent_dist
|
||||
return out.sample() if sample_posterior else out.mode()
|
||||
|
||||
|
||||
EntryClass = SAAudioVAEModel
|
||||
@@ -22,18 +22,23 @@ class Cosmos2VideoToWorldPipeline(ComposedPipelineBase):
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler", "safety_checker"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
use_karras_sigmas=True,
|
||||
)
|
||||
scheduler.config.sigma_max = 80.0
|
||||
scheduler.config.sigma_min = 0.002
|
||||
scheduler.config.sigma_data = 1.0
|
||||
scheduler.config.final_sigmas_type = "sigma_min"
|
||||
scheduler.sigma_max = 80.0
|
||||
scheduler.sigma_min = 0.002
|
||||
scheduler.sigma_data = 1.0
|
||||
self.modules["scheduler"] = scheduler
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
use_karras_sigmas=True)
|
||||
|
||||
sigma_max = 80.0
|
||||
sigma_min = 0.002
|
||||
sigma_data = 1.0
|
||||
final_sigmas_type = "sigma_min"
|
||||
|
||||
if self.modules["scheduler"] is not None:
|
||||
scheduler = self.modules["scheduler"]
|
||||
scheduler.config.sigma_max = sigma_max
|
||||
scheduler.config.sigma_min = sigma_min
|
||||
scheduler.config.sigma_data = sigma_data
|
||||
scheduler.config.final_sigmas_type = final_sigmas_type
|
||||
scheduler.sigma_max = sigma_max
|
||||
scheduler.sigma_min = sigma_min
|
||||
scheduler.sigma_data = sigma_data
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Importing continuation registers the "ltx2.v1" continuation kind with
|
||||
# the public compat layer so GenerationRequest.state(kind="ltx2.v1") is
|
||||
# recognized on the public API boundary.
|
||||
from fastvideo.pipelines.basic.ltx2 import continuation # noqa: F401
|
||||
|
||||
@@ -1,386 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Typed continuation state for the LTX-2 streaming pipeline.
|
||||
|
||||
Segment N+1 conditions on segment N's trailing decoded frames and
|
||||
denoised audio latents. The streaming runtime used to hold this state as
|
||||
per-worker globals; lifting it into a typed, JSON-serializable object
|
||||
lets clients snapshot, migrate, or round-trip it through an HTTP/RPC
|
||||
boundary. The envelope ``ContinuationState(kind, payload)`` is the
|
||||
shared public API; the typed class here owns the LTX-2 payload shape.
|
||||
|
||||
Serialization contract:
|
||||
|
||||
* Video frames → PNG bytes + base64, or a :class:`BlobStore` id.
|
||||
* Audio latents → a self-describing safetensors blob + base64, or a
|
||||
:class:`BlobStore` id. safetensors preserves ``bfloat16``, which a
|
||||
raw-numpy round-trip cannot.
|
||||
* The returned payload is always a plain JSON-serializable dict.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastvideo.api.compat import register_continuation_kind
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.entrypoints.streaming.session_store import BlobStore
|
||||
|
||||
LTX2_CONTINUATION_KIND = "ltx2.v1"
|
||||
"""Public ``ContinuationState.kind`` for LTX-2 payloads."""
|
||||
|
||||
LTX2_CONTINUATION_SCHEMA_VERSION = 1
|
||||
"""Payload schema version carried inside ``payload.schema_version``."""
|
||||
|
||||
DEFAULT_INLINE_THRESHOLD_BYTES = 2 * 1024 * 1024
|
||||
"""Tensors larger than this go to the blob store (if available). 2 MiB
|
||||
is below typical single-JSON-message limits (Dynamo: 4 MiB, Postgres
|
||||
TOAST: 1 GiB) and well above per-frame PNG payloads (~200 KiB at
|
||||
512x512)."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2ContinuationState:
|
||||
"""Typed LTX-2 continuation state carried between streaming segments.
|
||||
|
||||
``video_frames`` hold trailing decoded RGB frames (uint8 HxWx3) from
|
||||
segment N for conditioning segment N+1 via the VAE encode path.
|
||||
``audio_latents`` is the cached denoised audio latent tensor of shape
|
||||
``[B, C, T, mel]`` that segment N+1 will copy into the overlap
|
||||
region of its clean-latent conditioning.
|
||||
|
||||
Most fields map 1:1 onto the internal gpu_pool's per-worker state;
|
||||
the only new concept is the ``*_blob_id`` fields, which allow large
|
||||
tensors to live outside the JSON payload. See module docstring.
|
||||
"""
|
||||
|
||||
segment_index: int = 0
|
||||
"""Index of the *just-completed* segment. Segment 0 has no history;
|
||||
state returned after segment 0 carries ``segment_index=0`` and the
|
||||
caller uses ``segment_index + 1`` as the next segment number."""
|
||||
|
||||
video_frames: list[np.ndarray] | None = None
|
||||
"""Trailing decoded frames, each an RGB uint8 ``np.ndarray`` shaped
|
||||
``(H, W, 3)``. ``None`` when the state is blob-backed or unset."""
|
||||
|
||||
video_frames_blob_id: str | None = None
|
||||
"""Blob store id when the frames live outside the payload."""
|
||||
|
||||
video_conditioning_frame_idx: int = 0
|
||||
"""Target frame index inside the next segment that the trailing
|
||||
frames align with (matches the LTX-2 ``ltx2_video_conditions``
|
||||
tuple's ``frame_idx`` slot)."""
|
||||
|
||||
video_conditioning_strength: float = 1.0
|
||||
"""Conditioning strength in [0, 1]. Matches the ``ltx2_video_
|
||||
conditions`` tuple's strength slot."""
|
||||
|
||||
audio_latents: torch.Tensor | None = None
|
||||
"""Denoised audio latent tensor of shape ``[B, C, T, mel]``.
|
||||
``None`` when the state is blob-backed or unset."""
|
||||
|
||||
audio_latents_blob_id: str | None = None
|
||||
"""Blob store id when audio latents live outside the payload."""
|
||||
|
||||
audio_sample_rate: int | None = None
|
||||
"""Sample rate for the audio side (e.g. 24000)."""
|
||||
|
||||
audio_conditioning_num_frames: int = 0
|
||||
"""Number of trailing audio frames that carry over as clean context
|
||||
into segment N+1."""
|
||||
|
||||
audio_conditioning_strength: float = 1.0
|
||||
"""Clean-latent mask value applied to the overlap region; 0.0 keeps
|
||||
the cached audio entirely, 1.0 renoises from scratch."""
|
||||
|
||||
video_position_offset_sec: float = 0.0
|
||||
"""Seconds by which video RoPE is shifted forward so the audio
|
||||
prefix can sit at ``t >= 0`` when audio conditioning is longer than
|
||||
video conditioning."""
|
||||
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
"""Opaque metadata bag for forward-compat fields that don't need
|
||||
their own typed slot yet (e.g. custom knob experiments)."""
|
||||
|
||||
def to_continuation_state(
|
||||
self,
|
||||
*,
|
||||
blob_store: BlobStore | None = None,
|
||||
inline_threshold_bytes: int = DEFAULT_INLINE_THRESHOLD_BYTES,
|
||||
) -> ContinuationState:
|
||||
"""Serialize into a public :class:`ContinuationState`.
|
||||
|
||||
When ``blob_store`` is given, tensors larger than
|
||||
``inline_threshold_bytes`` are stored via
|
||||
:meth:`BlobStore.put` and referenced by id; otherwise all data
|
||||
is base64-encoded inline. The payload is always a plain
|
||||
JSON-serializable dict.
|
||||
"""
|
||||
payload: dict[str, Any] = {
|
||||
"schema_version": LTX2_CONTINUATION_SCHEMA_VERSION,
|
||||
"segment_index": int(self.segment_index),
|
||||
"video_conditioning_frame_idx": int(self.video_conditioning_frame_idx),
|
||||
"video_conditioning_strength": float(self.video_conditioning_strength),
|
||||
"audio_conditioning_num_frames": int(self.audio_conditioning_num_frames),
|
||||
"audio_conditioning_strength": float(self.audio_conditioning_strength),
|
||||
"video_position_offset_sec": float(self.video_position_offset_sec),
|
||||
"metadata": dict(self.metadata),
|
||||
}
|
||||
if self.audio_sample_rate is not None:
|
||||
payload["audio_sample_rate"] = int(self.audio_sample_rate)
|
||||
|
||||
video_payload = self._encode_video_frames(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=inline_threshold_bytes,
|
||||
)
|
||||
if video_payload is not None:
|
||||
payload["video"] = video_payload
|
||||
|
||||
audio_payload = self._encode_audio_latents(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=inline_threshold_bytes,
|
||||
)
|
||||
if audio_payload is not None:
|
||||
payload["audio"] = audio_payload
|
||||
|
||||
return ContinuationState(
|
||||
kind=LTX2_CONTINUATION_KIND,
|
||||
payload=payload,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_continuation_state(
|
||||
cls,
|
||||
state: ContinuationState,
|
||||
*,
|
||||
blob_store: BlobStore | None = None,
|
||||
) -> LTX2ContinuationState:
|
||||
"""Rebuild a typed state from a public :class:`ContinuationState`.
|
||||
|
||||
Raises :class:`ValueError` when the kind doesn't match or the
|
||||
schema version is unsupported.
|
||||
"""
|
||||
if state.kind != LTX2_CONTINUATION_KIND:
|
||||
raise ValueError(f"Expected ContinuationState.kind={LTX2_CONTINUATION_KIND!r}, "
|
||||
f"got {state.kind!r}")
|
||||
payload = state.payload or {}
|
||||
version = int(payload.get("schema_version", LTX2_CONTINUATION_SCHEMA_VERSION))
|
||||
if version != LTX2_CONTINUATION_SCHEMA_VERSION:
|
||||
raise ValueError(f"Unsupported LTX-2 continuation schema_version={version}; "
|
||||
f"this build expects {LTX2_CONTINUATION_SCHEMA_VERSION}")
|
||||
|
||||
out = cls(
|
||||
segment_index=int(payload.get("segment_index", 0)),
|
||||
video_conditioning_frame_idx=int(payload.get("video_conditioning_frame_idx", 0)),
|
||||
video_conditioning_strength=float(payload.get("video_conditioning_strength", 1.0)),
|
||||
audio_sample_rate=(int(payload["audio_sample_rate"]) if "audio_sample_rate" in payload else None),
|
||||
audio_conditioning_num_frames=int(payload.get("audio_conditioning_num_frames", 0)),
|
||||
audio_conditioning_strength=float(payload.get("audio_conditioning_strength", 1.0)),
|
||||
video_position_offset_sec=float(payload.get("video_position_offset_sec", 0.0)),
|
||||
metadata=dict(payload.get("metadata") or {}),
|
||||
)
|
||||
|
||||
video = payload.get("video")
|
||||
if isinstance(video, Mapping):
|
||||
cls._decode_video_frames(out, video, blob_store=blob_store)
|
||||
|
||||
audio = payload.get("audio")
|
||||
if isinstance(audio, Mapping):
|
||||
cls._decode_audio_latents(out, audio, blob_store=blob_store)
|
||||
|
||||
return out
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Video frame helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _encode_video_frames(
|
||||
self,
|
||||
*,
|
||||
blob_store: BlobStore | None,
|
||||
inline_threshold_bytes: int,
|
||||
) -> dict[str, Any] | None:
|
||||
if self.video_frames_blob_id is not None:
|
||||
return {"blob_id": self.video_frames_blob_id}
|
||||
if not self.video_frames:
|
||||
return None
|
||||
|
||||
encoded = [_encode_png(frame) for frame in self.video_frames]
|
||||
total = sum(len(b) for b in encoded)
|
||||
if blob_store is not None and total > inline_threshold_bytes:
|
||||
concatenated = _pack_frame_blobs(encoded)
|
||||
blob_id = blob_store.put(
|
||||
concatenated,
|
||||
mime="application/x-fastvideo-frames+png",
|
||||
)
|
||||
return {"blob_id": blob_id, "frame_count": len(encoded)}
|
||||
return {
|
||||
"frames_b64": [base64.b64encode(b).decode("ascii") for b in encoded],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _decode_video_frames(
|
||||
out: LTX2ContinuationState,
|
||||
video: Mapping[str, Any],
|
||||
*,
|
||||
blob_store: BlobStore | None,
|
||||
) -> None:
|
||||
blob_id = video.get("blob_id")
|
||||
if isinstance(blob_id, str):
|
||||
if blob_store is None:
|
||||
out.video_frames_blob_id = blob_id
|
||||
return
|
||||
raw = blob_store.get(blob_id)
|
||||
encoded = _unpack_frame_blobs(raw)
|
||||
out.video_frames = [_decode_png(b) for b in encoded]
|
||||
return
|
||||
frames_b64 = video.get("frames_b64")
|
||||
if isinstance(frames_b64, list):
|
||||
decoded = [_decode_png(base64.b64decode(b)) for b in frames_b64 if isinstance(b, str)]
|
||||
out.video_frames = decoded or None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Audio latent helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _encode_audio_latents(
|
||||
self,
|
||||
*,
|
||||
blob_store: BlobStore | None,
|
||||
inline_threshold_bytes: int,
|
||||
) -> dict[str, Any] | None:
|
||||
if self.audio_latents_blob_id is not None:
|
||||
return {"blob_id": self.audio_latents_blob_id}
|
||||
if self.audio_latents is None:
|
||||
return None
|
||||
raw = _tensor_to_safetensors_bytes(self.audio_latents)
|
||||
if blob_store is not None and len(raw) > inline_threshold_bytes:
|
||||
blob_id = blob_store.put(
|
||||
raw,
|
||||
mime="application/x-fastvideo-tensor+safetensors",
|
||||
)
|
||||
return {"blob_id": blob_id}
|
||||
return {"safetensors_b64": base64.b64encode(raw).decode("ascii")}
|
||||
|
||||
@staticmethod
|
||||
def _decode_audio_latents(
|
||||
out: LTX2ContinuationState,
|
||||
audio: Mapping[str, Any],
|
||||
*,
|
||||
blob_store: BlobStore | None,
|
||||
) -> None:
|
||||
blob_id = audio.get("blob_id")
|
||||
if isinstance(blob_id, str):
|
||||
if blob_store is None:
|
||||
out.audio_latents_blob_id = blob_id
|
||||
return
|
||||
raw = blob_store.get(blob_id)
|
||||
out.audio_latents = _safetensors_bytes_to_tensor(raw)
|
||||
return
|
||||
data_b64 = audio.get("safetensors_b64")
|
||||
if isinstance(data_b64, str):
|
||||
out.audio_latents = _safetensors_bytes_to_tensor(base64.b64decode(data_b64))
|
||||
|
||||
|
||||
def _encode_png(frame: np.ndarray) -> bytes:
|
||||
"""Encode an ``(H, W, 3)`` uint8 RGB frame as PNG bytes."""
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
if not isinstance(frame, np.ndarray):
|
||||
raise TypeError(f"LTX2 continuation frame must be a numpy ndarray, got {type(frame).__name__}")
|
||||
if frame.dtype != np.uint8 or frame.ndim != 3 or frame.shape[-1] != 3:
|
||||
raise ValueError("LTX2 continuation frame must be uint8 HxWx3 RGB; got "
|
||||
f"dtype={frame.dtype}, shape={frame.shape}")
|
||||
import io
|
||||
|
||||
buffer = io.BytesIO()
|
||||
Image.fromarray(frame).save(buffer, format="PNG")
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
def _decode_png(data: bytes) -> np.ndarray:
|
||||
import io
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
img = Image.open(io.BytesIO(data)).convert("RGB")
|
||||
return np.array(img, dtype=np.uint8)
|
||||
|
||||
|
||||
def _pack_frame_blobs(encoded: list[bytes]) -> bytes:
|
||||
"""Pack multiple PNG blobs into a single blob for blob-store storage.
|
||||
|
||||
Format: ``[4-byte big-endian count][4-byte len][png][4-byte len][png]...``.
|
||||
"""
|
||||
parts: list[bytes] = [len(encoded).to_bytes(4, "big")]
|
||||
for blob in encoded:
|
||||
parts.append(len(blob).to_bytes(4, "big"))
|
||||
parts.append(blob)
|
||||
return b"".join(parts)
|
||||
|
||||
|
||||
def _unpack_frame_blobs(raw: bytes) -> list[bytes]:
|
||||
if len(raw) < 4:
|
||||
raise ValueError("frame blob truncated: missing count header")
|
||||
count = int.from_bytes(raw[:4], "big")
|
||||
# Each frame contributes at least a 4-byte length prefix, so a
|
||||
# declared count larger than (len(raw) - 4) // 4 cannot fit and
|
||||
# would otherwise cause an O(count) allocation loop on malformed
|
||||
# input.
|
||||
if count > (len(raw) - 4) // 4:
|
||||
raise ValueError(f"frame blob declares {count} frames but buffer holds at most "
|
||||
f"{(len(raw) - 4) // 4}")
|
||||
out: list[bytes] = []
|
||||
cursor = 4
|
||||
for index in range(count):
|
||||
if cursor + 4 > len(raw):
|
||||
raise ValueError(f"frame blob truncated at frame {index} length header")
|
||||
length = int.from_bytes(raw[cursor:cursor + 4], "big")
|
||||
cursor += 4
|
||||
if cursor + length > len(raw):
|
||||
raise ValueError(f"frame blob truncated at frame {index} payload")
|
||||
out.append(raw[cursor:cursor + length])
|
||||
cursor += length
|
||||
return out
|
||||
|
||||
|
||||
def _tensor_to_safetensors_bytes(tensor: Any) -> bytes:
|
||||
"""Serialize a torch tensor to a self-describing safetensors blob.
|
||||
|
||||
Uses the in-memory safetensors API so the wire format preserves
|
||||
dtype (including ``bfloat16``, which a raw-numpy path cannot) and
|
||||
shape without needing sidecar metadata.
|
||||
"""
|
||||
import torch
|
||||
from safetensors.torch import save as st_save
|
||||
|
||||
if isinstance(tensor, torch.Tensor):
|
||||
return st_save({"t": tensor.detach().cpu()})
|
||||
import numpy as np
|
||||
if isinstance(tensor, np.ndarray):
|
||||
return st_save({"t": torch.from_numpy(np.ascontiguousarray(tensor))})
|
||||
raise TypeError("LTX2 audio_latents must be a torch.Tensor or numpy.ndarray, got "
|
||||
f"{type(tensor).__name__}")
|
||||
|
||||
|
||||
def _safetensors_bytes_to_tensor(raw: bytes) -> Any:
|
||||
from safetensors.torch import load as st_load
|
||||
return st_load(raw)["t"]
|
||||
|
||||
|
||||
register_continuation_kind(LTX2_CONTINUATION_KIND)
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_INLINE_THRESHOLD_BYTES",
|
||||
"LTX2ContinuationState",
|
||||
"LTX2_CONTINUATION_KIND",
|
||||
"LTX2_CONTINUATION_SCHEMA_VERSION",
|
||||
]
|
||||
@@ -2,7 +2,7 @@
|
||||
"""LTX2 model family pipeline presets."""
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
refine_stage_override_fields, )
|
||||
REFINE_STAGE_OVERRIDE_FIELDS, )
|
||||
|
||||
_LTX2_NEGATIVE_PROMPT = ("blurry, out of focus, overexposed, underexposed, low contrast, "
|
||||
"washed out colors, excessive noise, grainy texture, poor lighting, "
|
||||
@@ -36,7 +36,7 @@ _REFINE_STAGE = PresetStageSpec(
|
||||
name="refine",
|
||||
kind="refinement",
|
||||
description="Latent-upsample + second-pass refine",
|
||||
allowed_overrides=refine_stage_override_fields(),
|
||||
allowed_overrides=REFINE_STAGE_OVERRIDE_FIELDS,
|
||||
)
|
||||
|
||||
LTX2_BASE = InferencePreset(
|
||||
|
||||
@@ -27,8 +27,6 @@ class LTX2RefinePresetOverride:
|
||||
class LTX2RefineStageOverride:
|
||||
"""Per-request refine tuning under ``stage_overrides.refine``."""
|
||||
|
||||
# Stage-2 refine only validates 2 (reduced) and 3 (official distilled)
|
||||
# sigma schedules; other values raise at pipeline construction.
|
||||
num_inference_steps: int | None = None
|
||||
guidance_scale: float | None = None
|
||||
image_crf: int | None = None
|
||||
@@ -42,18 +40,15 @@ def refine_override_to_dict(override: LTX2RefinePresetOverride | LTX2RefineStage
|
||||
return {k: v for k, v in asdict(override).items() if v is not None}
|
||||
|
||||
|
||||
def refine_preset_override_fields() -> frozenset[str]:
|
||||
return frozenset(f.name for f in fields(LTX2RefinePresetOverride))
|
||||
|
||||
|
||||
def refine_stage_override_fields() -> frozenset[str]:
|
||||
return frozenset(f.name for f in fields(LTX2RefineStageOverride))
|
||||
|
||||
REFINE_PRESET_OVERRIDE_FIELDS: frozenset[str] = frozenset(f.name for f in fields(LTX2RefinePresetOverride))
|
||||
REFINE_STAGE_OVERRIDE_FIELDS: frozenset[str] = frozenset(f.name for f in fields(LTX2RefineStageOverride))
|
||||
REFINE_FLAT_KEYS: frozenset[str] = (REFINE_PRESET_OVERRIDE_FIELDS | REFINE_STAGE_OVERRIDE_FIELDS)
|
||||
|
||||
__all__ = [
|
||||
"LTX2RefinePresetOverride",
|
||||
"LTX2RefineStageOverride",
|
||||
"REFINE_FLAT_KEYS",
|
||||
"REFINE_PRESET_OVERRIDE_FIELDS",
|
||||
"REFINE_STAGE_OVERRIDE_FIELDS",
|
||||
"refine_override_to_dict",
|
||||
"refine_preset_override_fields",
|
||||
"refine_stage_override_fields",
|
||||
]
|
||||
|
||||
@@ -26,7 +26,7 @@ MATRIXGAME_I2V = InferencePreset(
|
||||
"fps": 25,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 3,
|
||||
"negative_prompt": "",
|
||||
"negative_prompt": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -1,63 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio presets.
|
||||
|
||||
Sampling defaults track the published HF model card
|
||||
(https://huggingface.co/stabilityai/stable-audio-open-1.0):
|
||||
100 steps, CFG=7, dpmpp-3m-sde, sigma_min=0.3, sigma_max=500, rho=1.0.
|
||||
"""
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Stable Audio Cosine-DPM++ denoising with text + duration CFG.",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
# `audio_start_in_s` / `audio_end_in_s` are call-kwargs, kept off here.
|
||||
# `height`/`width` are pinned to 8 (the shared `InputValidationStage`
|
||||
# rejects values that aren't divisible by 8) and `num_frames` to 1 so
|
||||
# the video-shaped preallocation in `VideoGenerator` stays tiny — the
|
||||
# real output is the audio waveform on `result["audio"]`, not the
|
||||
# placeholder frame tensor.
|
||||
_SHARED_DEFAULTS = {
|
||||
"seed": 0,
|
||||
"guidance_scale": 7.0,
|
||||
"num_inference_steps": 100,
|
||||
"negative_prompt": "",
|
||||
"height": 8,
|
||||
"width": 8,
|
||||
"num_frames": 1,
|
||||
}
|
||||
|
||||
STABLE_AUDIO_OPEN_1_0_BASE = InferencePreset(
|
||||
name="stable_audio_open_1_0_base",
|
||||
version=1,
|
||||
model_family="stable_audio",
|
||||
description=("Stability AI Stable Audio Open 1.0 text-to-audio. Generates up "
|
||||
"to ~47.5s of stereo 44.1 kHz audio per call. Default duration "
|
||||
"is 10s; raise via `audio_end_in_s` up to the model max."),
|
||||
workload_type="t2v", # NOTE: WorkloadType has no T2A variant yet (REVIEW item 28)
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults=dict(_SHARED_DEFAULTS),
|
||||
)
|
||||
|
||||
# Smaller / faster checkpoint with the same Oobleck VAE but a 1024-dim
|
||||
# 16-layer DiT with `qk_norm="ln"`. Sampling defaults match the official
|
||||
# `stable-audio-open-small` model card.
|
||||
STABLE_AUDIO_OPEN_SMALL = InferencePreset(
|
||||
name="stable_audio_open_small",
|
||||
version=1,
|
||||
model_family="stable_audio",
|
||||
description=("Stability AI Stable Audio Open Small text-to-audio. Faster than "
|
||||
"the 1.0 base; supports up to ~11.9s of stereo 44.1 kHz audio per "
|
||||
"call (smaller training window)."),
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults=dict(_SHARED_DEFAULTS),
|
||||
)
|
||||
|
||||
ALL_PRESETS = (STABLE_AUDIO_OPEN_1_0_BASE, STABLE_AUDIO_OPEN_SMALL)
|
||||
@@ -1,125 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 pipeline (T2A + A2A + RePaint inpainting).
|
||||
|
||||
Components are loaded via the standard
|
||||
`ComposedPipelineBase.load_modules` against the FastVideo-curated
|
||||
Diffusers-format repo `FastVideo/stable-audio-open-1.0-Diffusers`
|
||||
(produced by
|
||||
`scripts/checkpoint_conversion/stable_audio_to_diffusers.py`). The DiT
|
||||
is a `BaseDiT` subclass loaded by `TransformerLoader`; the VAE is
|
||||
loaded by `VAELoader`; the multi-conditioner (T5 + NumberConditioners)
|
||||
is loaded by `ConditionerLoader` (a Stable Audio-specific addition).
|
||||
|
||||
Stages:
|
||||
|
||||
InputValidationStage
|
||||
→ StableAudioConditioningStage (T5 + NumberConditioner -> cross-attn + global cond, with CFG)
|
||||
→ StableAudioLatentPreparationStage (initial Gaussian noise; encodes A2A / inpaint refs)
|
||||
→ StableAudioDenoisingStage (k-diffusion `dpmpp-3m-sde` over the DiT)
|
||||
→ StableAudioDecodingStage (OobleckVAE -> waveform)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.basic.stable_audio.stages import (
|
||||
StableAudioConditioningStage,
|
||||
StableAudioDecodingStage,
|
||||
StableAudioDenoisingStage,
|
||||
StableAudioLatentPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import InputValidationStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _warn_tf32_disabled_for_stable_audio() -> None:
|
||||
logger.warning("Stable Audio pipeline is disabling process-global "
|
||||
"torch.backends.{cuda.matmul.allow_tf32, cudnn.allow_tf32, "
|
||||
"cuda.matmul.allow_fp16_reduced_precision_reduction, "
|
||||
"cudnn.benchmark} for A2A renoise determinism. Other models "
|
||||
"loaded into this process will inherit these settings.")
|
||||
|
||||
|
||||
def _disable_tf32_for_stable_audio() -> None:
|
||||
"""Disable TF32 / cuDNN nondeterminism — A2A renoise-then-denoise SDE
|
||||
amplifies per-element drift, and the published parity bounds were
|
||||
set with these off. Process-global; the first call logs a warning.
|
||||
"""
|
||||
_warn_tf32_disabled_for_stable_audio()
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction = False
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
class StableAudioPipeline(ComposedPipelineBase):
|
||||
"""Stable Audio Open 1.0 pipeline.
|
||||
|
||||
Mode is kwargs-driven on `generate_video()`:
|
||||
|
||||
* Text-to-audio (default) -- `prompt=...`, `audio_end_in_s=...`
|
||||
* Audio-to-audio variation -- add `init_audio=ref` (and optionally
|
||||
`init_noise_level`, lower = closer to reference)
|
||||
* RePaint inpainting / outpainting -- add `inpaint_audio=ref` and
|
||||
`inpaint_mask` (1-D, 1 = keep / 0 = regenerate)
|
||||
|
||||
See `examples/inference/basic/basic_stable_audio*.py` for runnable
|
||||
examples of each mode.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"vae",
|
||||
"transformer",
|
||||
"conditioner",
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Apply Stable Audio's process-global numerics overrides BEFORE
|
||||
the standard component loaders run (TF32 off for A2A renoise
|
||||
determinism)."""
|
||||
_disable_tf32_for_stable_audio()
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="conditioning_stage",
|
||||
stage=StableAudioConditioningStage(conditioner=self.get_module("conditioner")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=StableAudioLatentPreparationStage(
|
||||
io_channels=64,
|
||||
# Per-variant training window: 2,097,152 (~47.5s) for
|
||||
# SA-1.0; 524,288 (~11.9s) for SA-small. Pulled from the
|
||||
# pipeline config so each variant gets its own latent
|
||||
# length.
|
||||
sample_size=pc.sample_size,
|
||||
vae=self.get_module("vae"),
|
||||
sample_rate=pc.sampling_rate,
|
||||
audio_channels=pc.audio_channels,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=StableAudioDenoisingStage(transformer=self.get_module("transformer")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=StableAudioDecodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = StableAudioPipeline
|
||||
@@ -1,12 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.pipelines.basic.stable_audio.stages.conditioning import StableAudioConditioningStage
|
||||
from fastvideo.pipelines.basic.stable_audio.stages.decoding import StableAudioDecodingStage
|
||||
from fastvideo.pipelines.basic.stable_audio.stages.denoising import StableAudioDenoisingStage
|
||||
from fastvideo.pipelines.basic.stable_audio.stages.latent_preparation import StableAudioLatentPreparationStage
|
||||
|
||||
__all__ = [
|
||||
"StableAudioConditioningStage",
|
||||
"StableAudioDecodingStage",
|
||||
"StableAudioDenoisingStage",
|
||||
"StableAudioLatentPreparationStage",
|
||||
]
|
||||
@@ -1,98 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio conditioning stage."""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class StableAudioConditioningStage(PipelineStage):
|
||||
"""Run the conditioner over the prompt + duration and stash the
|
||||
DiT-ready (cross_attn_cond, cross_attn_mask, global_embed) triple
|
||||
on `batch.extra` (plus the negative-prompt triple when CFG is on).
|
||||
"""
|
||||
|
||||
def __init__(self, conditioner) -> None:
|
||||
super().__init__()
|
||||
self.conditioner = conditioner
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
device = next(self.conditioner.parameters()).device
|
||||
|
||||
start_attr = getattr(batch, "audio_start_in_s", None)
|
||||
end_attr = getattr(batch, "audio_end_in_s", None)
|
||||
audio_start_in_s = float(start_attr if start_attr is not None else pc.audio_start_in_s)
|
||||
audio_end_in_s = float(end_attr if end_attr is not None else pc.audio_end_in_s)
|
||||
max_duration = float(getattr(pc, "max_audio_duration_s", 2097152 / 44100))
|
||||
if audio_start_in_s < 0:
|
||||
raise ValueError(f"audio_start_in_s must be >= 0, got {audio_start_in_s}.")
|
||||
if audio_end_in_s <= audio_start_in_s:
|
||||
raise ValueError(f"audio_end_in_s ({audio_end_in_s}) must be > audio_start_in_s "
|
||||
f"({audio_start_in_s}).")
|
||||
if audio_end_in_s > max_duration:
|
||||
raise ValueError(f"audio_end_in_s ({audio_end_in_s}s) exceeds the model's fixed "
|
||||
f"window of {max_duration:.4f}s. Stable Audio Open 1.0 always "
|
||||
f"samples a 2,097,152-frame latent and slices to "
|
||||
f"[start, end] after decode; values past the window are silently "
|
||||
f"truncated. Lower audio_end_in_s or split the request.")
|
||||
guidance_scale = float(batch.guidance_scale or pc.guidance_scale)
|
||||
do_cfg = guidance_scale > 1.0
|
||||
|
||||
if isinstance(batch.prompt, str):
|
||||
prompt = batch.prompt
|
||||
elif isinstance(batch.prompt, list):
|
||||
if len(batch.prompt) > 1:
|
||||
raise ValueError(f"Stable Audio does not support batched prompts; got "
|
||||
f"{len(batch.prompt)} entries. Pass a single string or a "
|
||||
f"single-element list.")
|
||||
prompt = batch.prompt[0] if batch.prompt else ""
|
||||
else:
|
||||
raise TypeError(f"`prompt` must be a string or a list of strings, got "
|
||||
f"{type(batch.prompt).__name__}.")
|
||||
# Send only the keys the conditioner declares (per-variant).
|
||||
all_cond_values = {
|
||||
"prompt": prompt,
|
||||
"seconds_start": audio_start_in_s,
|
||||
"seconds_total": audio_end_in_s,
|
||||
}
|
||||
active_ids = self.conditioner.cross_attention_cond_ids
|
||||
cond_meta = [{k: all_cond_values[k] for k in active_ids if k in all_cond_values}]
|
||||
cond = self.conditioner(cond_meta, device)
|
||||
cross_attn_cond, cross_attn_mask, global_embed = self.conditioner.get_conditioning_inputs(cond)
|
||||
|
||||
neg_cross_attn_cond = None
|
||||
neg_cross_attn_mask = None
|
||||
neg_global_embed = None
|
||||
if do_cfg:
|
||||
neg_prompt = batch.negative_prompt or ""
|
||||
if isinstance(neg_prompt, list):
|
||||
neg_prompt = neg_prompt[0] if neg_prompt else ""
|
||||
neg_values = dict(all_cond_values, prompt=neg_prompt)
|
||||
neg_meta = [{k: neg_values[k] for k in active_ids if k in neg_values}]
|
||||
neg = self.conditioner(neg_meta, device)
|
||||
neg_cross_attn_cond, neg_cross_attn_mask, neg_global_embed = (self.conditioner.get_conditioning_inputs(neg))
|
||||
|
||||
if batch.extra is None:
|
||||
batch.extra = {}
|
||||
batch.extra["cross_attn_cond"] = cross_attn_cond
|
||||
batch.extra["cross_attn_mask"] = cross_attn_mask
|
||||
batch.extra["global_embed"] = global_embed
|
||||
batch.extra["negative_cross_attn_cond"] = neg_cross_attn_cond
|
||||
batch.extra["negative_cross_attn_mask"] = neg_cross_attn_mask
|
||||
batch.extra["negative_global_embed"] = neg_global_embed
|
||||
batch.extra["do_cfg"] = do_cfg
|
||||
batch.extra["audio_start_in_s"] = audio_start_in_s
|
||||
batch.extra["audio_end_in_s"] = audio_end_in_s
|
||||
return batch
|
||||
@@ -1,67 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio decoding: latent -> waveform via OobleckVAE.
|
||||
|
||||
Slices the output to `[audio_start_in_s, audio_end_in_s]` and stashes
|
||||
the result on `batch.extra["audio"]` + `["audio_sample_rate"]` for
|
||||
`VideoGenerator._mux_audio` to pick up.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class StableAudioDecodingStage(PipelineStage):
|
||||
"""Decode latent → audio waveform + slice to [start, end]."""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
latents = batch.latents
|
||||
|
||||
# VAE may be CPU-parked under `vae_cpu_offload=True`.
|
||||
from fastvideo.distributed.parallel_state import get_local_torch_device
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
decoded = self.vae.decode(latents)
|
||||
if hasattr(decoded, "sample"): # tolerate tensor or dataclass
|
||||
decoded = decoded.sample
|
||||
|
||||
sr = int(getattr(self.vae, "sampling_rate", pc.sampling_rate))
|
||||
start_in_s = float(batch.extra.get("audio_start_in_s", pc.audio_start_in_s))
|
||||
end_in_s = float(batch.extra.get("audio_end_in_s", pc.audio_end_in_s))
|
||||
decoded = decoded[:, :, int(start_in_s * sr):int(end_in_s * sr)]
|
||||
|
||||
if batch.extra is None:
|
||||
batch.extra = {}
|
||||
# `_mux_audio` / `_write_pcm_wav` want `[samples, channels]`.
|
||||
batch.extra["audio"] = decoded.squeeze(0).T.detach().float().cpu().numpy()
|
||||
batch.extra["audio_sample_rate"] = sr
|
||||
batch.extra["audio_only"] = True
|
||||
# Raw tensor for parity tests.
|
||||
batch.extra["decoded_audio"] = decoded.detach().cpu()
|
||||
|
||||
# `VideoGenerator.generate_video` is video-shaped (asserts
|
||||
# `output_batch.output is not None`); fill with a placeholder of
|
||||
# the expected `[B, 3, num_frames, H, W]` shape — the real audio
|
||||
# is on `batch.extra` above. Pure-audio workload support tracked
|
||||
# in REVIEW item 28.
|
||||
b = decoded.shape[0]
|
||||
n_frames = int(getattr(batch, "num_frames", 1) or 1)
|
||||
h = int(getattr(batch, "height", 1) or 1)
|
||||
w = int(getattr(batch, "width", 1) or 1)
|
||||
batch.output = torch.zeros((b, 3, n_frames, h, w), dtype=torch.uint8)
|
||||
return batch
|
||||
@@ -1,207 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio denoising — k-diffusion `dpmpp-3m-sde` over the DiT.
|
||||
|
||||
CFG-batched conditioning is built once outside the sampler loop so the
|
||||
adapter only does `cat([x, x])` + DiT call per step.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class _DiTAdapter(nn.Module):
|
||||
"""`StableAudioDiT` -> `K.external.VDenoiser` adapter.
|
||||
|
||||
`batch_cond` / `batch_global` are precomputed CFG-batched tensors
|
||||
(`[2, ...]` for CFG, `[1, ...]` otherwise); building them once
|
||||
outside the sampler loop saves ~3 cats × 100 steps per call.
|
||||
"""
|
||||
|
||||
def __init__(self, dit, *, batch_cond: torch.Tensor, batch_global: torch.Tensor, cfg_scale: float) -> None:
|
||||
super().__init__()
|
||||
self.dit = dit
|
||||
self.batch_cond = batch_cond
|
||||
self.batch_global = batch_global
|
||||
self.cfg_scale = cfg_scale
|
||||
self.do_cfg = cfg_scale != 1.0
|
||||
|
||||
def forward(self, x: torch.Tensor, t: torch.Tensor, **_unused) -> torch.Tensor:
|
||||
if not self.do_cfg:
|
||||
return self.dit(x, t, cross_attn_cond=self.batch_cond, global_embed=self.batch_global)
|
||||
batch_x = torch.cat([x, x], dim=0)
|
||||
batch_t = torch.cat([t, t], dim=0)
|
||||
out = self.dit(batch_x, batch_t, cross_attn_cond=self.batch_cond, global_embed=self.batch_global)
|
||||
cond_out, uncond_out = torch.chunk(out, 2, dim=0)
|
||||
return uncond_out + (cond_out - uncond_out) * self.cfg_scale
|
||||
|
||||
|
||||
class StableAudioDenoisingStage(PipelineStage):
|
||||
"""k-diffusion `dpmpp-3m-sde` sampling loop."""
|
||||
|
||||
# Sampler defaults from the published model card.
|
||||
_SIGMA_MIN = 0.3
|
||||
_SIGMA_MAX = 500.0
|
||||
_RHO = 1.0
|
||||
_LOG_SIGMA_MIN = math.log(_SIGMA_MIN)
|
||||
_LOG_SIGMA_MAX = math.log(_SIGMA_MAX)
|
||||
|
||||
def __init__(self, transformer) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
|
||||
def _resolve_sigma_max(self, batch) -> float:
|
||||
"""Map A2A intent to `sigma_max`.
|
||||
|
||||
Public knob is `init_audio_strength` (0..1, higher = closer to
|
||||
source), log-interpolated between SIGMA_MIN (= preservation) and
|
||||
SIGMA_MAX (= full T2A). Raw `init_noise_level` is the legacy
|
||||
sigma_max override; passing both is an error.
|
||||
"""
|
||||
raw = getattr(batch, "init_noise_level", None)
|
||||
strength = getattr(batch, "init_audio_strength", None)
|
||||
if raw is not None and strength is not None:
|
||||
raise ValueError("Pass `init_audio_strength` (0..1) OR `init_noise_level` "
|
||||
"(raw sigma_max), not both.")
|
||||
if raw is not None:
|
||||
return float(raw)
|
||||
s = max(0.0, min(1.0, float(strength) if strength is not None else 0.6))
|
||||
return float(math.exp(self._LOG_SIGMA_MAX - s * (self._LOG_SIGMA_MAX - self._LOG_SIGMA_MIN)))
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
ext = batch.extra
|
||||
device = batch.latents.device
|
||||
guidance_scale = float(batch.guidance_scale or pc.guidance_scale)
|
||||
steps = int(batch.num_inference_steps)
|
||||
|
||||
import k_diffusion as K
|
||||
|
||||
init_latent = ext.get("init_latent")
|
||||
sigma_max = self._resolve_sigma_max(batch) if init_latent is not None else self._SIGMA_MAX
|
||||
|
||||
sigmas = K.sampling.get_sigmas_polyexponential(steps, self._SIGMA_MIN, sigma_max, self._RHO, device=device)
|
||||
# Cast noise + conditioning to the DiT's dtype before sampling
|
||||
# (matches `stable_audio_tools/inference/generation.py:185-187`).
|
||||
model_dtype = next(self.transformer.parameters()).dtype
|
||||
|
||||
def _cast(t: torch.Tensor | None) -> torch.Tensor | None:
|
||||
return t.to(model_dtype) if t is not None else None
|
||||
|
||||
x = (batch.latents * sigmas[0]).to(model_dtype)
|
||||
if init_latent is not None:
|
||||
x = x + init_latent.to(model_dtype)
|
||||
|
||||
batch_cond, batch_global = _build_cfg_conditioning(
|
||||
cross_attn_cond=ext["cross_attn_cond"].to(model_dtype),
|
||||
global_embed=ext["global_embed"].to(model_dtype),
|
||||
negative_cross_attn_cond=_cast(ext.get("negative_cross_attn_cond")),
|
||||
negative_cross_attn_mask=ext.get("negative_cross_attn_mask"),
|
||||
negative_global_embed=_cast(ext.get("negative_global_embed")),
|
||||
do_cfg=guidance_scale != 1.0,
|
||||
)
|
||||
adapter = _DiTAdapter(self.transformer,
|
||||
batch_cond=batch_cond,
|
||||
batch_global=batch_global,
|
||||
cfg_scale=guidance_scale)
|
||||
denoiser = K.external.VDenoiser(adapter)
|
||||
|
||||
# RePaint blending hook — works on any v-prediction model, no
|
||||
# inpaint-trained checkpoint needed.
|
||||
inpaint_mask = ext.get("inpaint_mask_latent")
|
||||
inpaint_ref = ext.get("inpaint_reference_latent")
|
||||
if inpaint_mask is not None and inpaint_ref is not None:
|
||||
inpaint_mask = inpaint_mask.to(model_dtype)
|
||||
inpaint_ref = inpaint_ref.to(model_dtype)
|
||||
callback = _make_inpaint_callback(inpaint_ref, inpaint_mask, sigmas)
|
||||
else:
|
||||
callback = None
|
||||
|
||||
# `LocalAttention` (in `StableAudioDiT`) reads `get_forward_context()`
|
||||
# for `attn_metadata`; wrap the whole loop.
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
sampled = K.sampling.sample_dpmpp_3m_sde(denoiser,
|
||||
x,
|
||||
sigmas,
|
||||
disable=False,
|
||||
extra_args={},
|
||||
callback=callback)
|
||||
|
||||
# Final blend so the kept region of the inpaint reference is exact.
|
||||
if inpaint_mask is not None and inpaint_ref is not None:
|
||||
sampled = inpaint_ref * inpaint_mask + sampled * (1 - inpaint_mask)
|
||||
batch.latents = sampled
|
||||
return batch
|
||||
|
||||
|
||||
def _build_cfg_conditioning(
|
||||
*,
|
||||
cross_attn_cond: torch.Tensor,
|
||||
global_embed: torch.Tensor,
|
||||
negative_cross_attn_cond: torch.Tensor | None,
|
||||
negative_cross_attn_mask: torch.Tensor | None,
|
||||
negative_global_embed: torch.Tensor | None,
|
||||
do_cfg: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Build the CFG-batched `(cond, global)` tensors once.
|
||||
|
||||
Cond ordering is `[conditioned, unconditioned]` (the adapter splits
|
||||
with the same convention). Masked negative cond is zero-filled where
|
||||
`mask == 0`.
|
||||
"""
|
||||
if not do_cfg:
|
||||
return cross_attn_cond, global_embed
|
||||
if negative_cross_attn_cond is not None:
|
||||
if negative_cross_attn_mask is not None:
|
||||
neg_mask = negative_cross_attn_mask.to(torch.bool).unsqueeze(2)
|
||||
null_embed = torch.zeros_like(cross_attn_cond)
|
||||
negative_cross_attn_cond = torch.where(neg_mask, negative_cross_attn_cond, null_embed)
|
||||
batch_cond = torch.cat([cross_attn_cond, negative_cross_attn_cond], dim=0)
|
||||
else:
|
||||
batch_cond = torch.cat([cross_attn_cond, torch.zeros_like(cross_attn_cond)], dim=0)
|
||||
other_global = global_embed if negative_global_embed is None else negative_global_embed
|
||||
batch_global = torch.cat([global_embed, other_global], dim=0)
|
||||
return batch_cond, batch_global
|
||||
|
||||
|
||||
def _make_inpaint_callback(reference_latent: torch.Tensor, mask: torch.Tensor, sigmas: torch.Tensor):
|
||||
"""RePaint blending callback for the k-diffusion sampler.
|
||||
|
||||
At every step, replaces the kept region (`mask == 1`) of the in-
|
||||
flight latent with the reference re-noised to the next sigma —
|
||||
pulls the kept region back onto the trajectory the model expects,
|
||||
so RePaint-style inpainting converges on non-inpaint-trained models.
|
||||
|
||||
Pre-allocates the noise buffer so the ~100 sampler steps don't churn
|
||||
~25 MB of fresh allocations per call.
|
||||
"""
|
||||
noise_buf = torch.empty_like(reference_latent)
|
||||
inv_mask = 1 - mask
|
||||
|
||||
def cb(info: dict) -> None:
|
||||
i = int(info["i"])
|
||||
next_i = min(i + 1, len(sigmas) - 1)
|
||||
sigma_next = float(sigmas[next_i])
|
||||
noise_buf.normal_()
|
||||
# `state["x"]` is the live latent; the dpmpp-3m-sde sampler picks
|
||||
# up our in-place mutation between steps (verified against
|
||||
# k_diffusion 0.1.1.post1).
|
||||
x = info["x"]
|
||||
x.copy_((reference_latent + noise_buf * sigma_next) * mask + x * inv_mask)
|
||||
|
||||
return cb
|
||||
@@ -1,174 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio latent preparation.
|
||||
|
||||
Seeds + samples the initial Gaussian noise; encodes `init_audio` (A2A
|
||||
variation) or `inpaint_audio` + `inpaint_mask` (RePaint inpainting) into
|
||||
latent-space tensors on `batch.extra` for the denoising stage.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class StableAudioLatentPreparationStage(PipelineStage):
|
||||
|
||||
def __init__(self,
|
||||
io_channels: int = 64,
|
||||
sample_size: int = 2097152,
|
||||
vae=None,
|
||||
sample_rate: int = 44100,
|
||||
audio_channels: int = 2) -> None:
|
||||
super().__init__()
|
||||
self.io_channels = io_channels
|
||||
# Audio-domain length the model was trained for; latent length
|
||||
# = sample_size // vae.hop_length (= 2097152 / 2048 = 1024).
|
||||
self.sample_size = sample_size
|
||||
self.vae = vae # used to encode init_audio / inpaint_audio
|
||||
self.sample_rate = sample_rate
|
||||
self.audio_channels = audio_channels
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
ext = batch.extra or {}
|
||||
device = ext["cross_attn_cond"].device
|
||||
latent_sample_size = self.sample_size // self._hop_length()
|
||||
|
||||
seed = int(batch.seed) if batch.seed is not None else 0
|
||||
torch.manual_seed(seed)
|
||||
latents = torch.randn((1, self.io_channels, latent_sample_size), device=device)
|
||||
|
||||
batch.latents = latents
|
||||
if batch.extra is None:
|
||||
batch.extra = {}
|
||||
|
||||
init_audio = getattr(batch, "init_audio", None)
|
||||
inpaint_audio = getattr(batch, "inpaint_audio", None)
|
||||
inpaint_mask = getattr(batch, "inpaint_mask", None)
|
||||
|
||||
# Loud-fail rather than silently falling through to T2A.
|
||||
if inpaint_audio is not None and inpaint_mask is None:
|
||||
raise ValueError("Stable Audio inpainting requires both `inpaint_audio` and "
|
||||
"`inpaint_mask` (1-D tensor in {0, 1} at the model sample rate, "
|
||||
"1 = keep, 0 = regenerate). Got `inpaint_audio` without `inpaint_mask`.")
|
||||
if inpaint_mask is not None and inpaint_audio is None:
|
||||
raise ValueError("Stable Audio inpainting requires both `inpaint_audio` and "
|
||||
"`inpaint_mask`. Got `inpaint_mask` without `inpaint_audio` — "
|
||||
"did you mean to pass `init_audio` (audio-to-audio variation)?")
|
||||
if init_audio is not None and inpaint_audio is not None:
|
||||
raise ValueError("Stable Audio cannot do A2A variation and inpainting in the "
|
||||
"same call. Pass either `init_audio` (variation) or "
|
||||
"`inpaint_audio` + `inpaint_mask` (inpainting), not both.")
|
||||
|
||||
if init_audio is not None:
|
||||
batch.extra["init_latent"] = self._encode_audio_reference(init_audio, device)
|
||||
|
||||
if inpaint_audio is not None and inpaint_mask is not None:
|
||||
batch.extra["inpaint_reference_latent"] = self._encode_audio_reference(inpaint_audio, device)
|
||||
batch.extra["inpaint_mask_latent"] = self._prepare_mask(inpaint_mask, latent_sample_size, device)
|
||||
return batch
|
||||
|
||||
def _hop_length(self) -> int:
|
||||
return int(self.vae.hop_length)
|
||||
|
||||
def _encode_audio_reference(self, audio, device: torch.device) -> torch.Tensor:
|
||||
"""Pad/truncate to `sample_size` and encode via the VAE.
|
||||
|
||||
`audio` may be a tensor (`[samples]`, `[C, samples]`, or
|
||||
`[B, C, samples]`) at the model's sample rate, or a path to any
|
||||
audio-bearing file (`.wav` / `.mp3` / `.mp4` / `.m4a` / `.flac`,
|
||||
...) that PyAV can decode — we resample on load so callers don't
|
||||
have to.
|
||||
"""
|
||||
assert self.vae is not None, "VAE required for init_audio / inpaint_audio encoding"
|
||||
# VAE may be CPU-parked under `vae_cpu_offload=True`.
|
||||
self.vae = self.vae.to(device)
|
||||
if isinstance(audio, str | os.PathLike):
|
||||
audio = _decode_audio_file(audio, target_sr=self.sample_rate)
|
||||
audio = audio.to(device=device, dtype=torch.float32)
|
||||
if audio.dim() == 1:
|
||||
audio = audio.unsqueeze(0).unsqueeze(0)
|
||||
elif audio.dim() == 2:
|
||||
audio = audio.unsqueeze(0)
|
||||
# Match expected channel count (mono → repeat to stereo).
|
||||
if audio.shape[1] == 1 and self.audio_channels == 2:
|
||||
audio = audio.repeat(1, 2, 1)
|
||||
elif audio.shape[1] == 2 and self.audio_channels == 1:
|
||||
audio = audio.mean(dim=1, keepdim=True)
|
||||
# Pad/truncate to model sample_size.
|
||||
cur_len = audio.shape[-1]
|
||||
if cur_len < self.sample_size:
|
||||
audio = F.pad(audio, (0, self.sample_size - cur_len))
|
||||
elif cur_len > self.sample_size:
|
||||
audio = audio[..., :self.sample_size]
|
||||
# Stochastic sample (the next random draw after the latent
|
||||
# `randn` above), so encode-noise stays on the seeded sequence.
|
||||
return self.vae.encode(audio.to(next(self.vae.parameters()).dtype)).sample()
|
||||
|
||||
def _prepare_mask(self, mask, latent_len: int, device: torch.device) -> torch.Tensor:
|
||||
"""Pad/truncate a binary mask to `sample_size`, then
|
||||
nearest-resample to `[1, 1, latent_len]`. Convention: 1 = keep
|
||||
the reference, 0 = regenerate.
|
||||
|
||||
`mask` may be a `[samples]` tensor at the model sample rate or a
|
||||
`(keep_seconds, total_seconds)` tuple — the tuple form builds
|
||||
"keep first K seconds, regenerate the rest" automatically.
|
||||
"""
|
||||
if isinstance(mask, tuple) and len(mask) == 2:
|
||||
keep_s, total_s = (float(x) for x in mask)
|
||||
keep_n = int(keep_s * self.sample_rate)
|
||||
total_n = int(total_s * self.sample_rate)
|
||||
mask = torch.zeros(total_n, dtype=torch.float32)
|
||||
mask[:keep_n] = 1.0
|
||||
m = mask.to(device=device, dtype=torch.float32)
|
||||
if m.dim() == 1:
|
||||
m = m.unsqueeze(0)
|
||||
cur_len = m.shape[-1]
|
||||
if cur_len < self.sample_size:
|
||||
m = F.pad(m, (0, self.sample_size - cur_len))
|
||||
elif cur_len > self.sample_size:
|
||||
m = m[..., :self.sample_size]
|
||||
return F.interpolate(m.unsqueeze(1), size=latent_len, mode="nearest")
|
||||
|
||||
|
||||
def _decode_audio_file(path, target_sr: int) -> torch.Tensor:
|
||||
"""Decode any audio-bearing file (wav, mp3, mp4, m4a, flac, ...) via
|
||||
PyAV and resample to `target_sr`. Returns `[channels, samples]`
|
||||
float32 in roughly [-1, 1].
|
||||
|
||||
PyAV is already a FastVideo dep (used for muxing in
|
||||
`VideoGenerator._mux_audio`). `torchaudio.load` on container
|
||||
formats (mp4 / m4a) routes through `torchcodec`, which pulls in a
|
||||
full CUDA NVRTC stack we don't otherwise need.
|
||||
"""
|
||||
import av
|
||||
import numpy as np
|
||||
container = av.open(str(path))
|
||||
audio_stream = next(s for s in container.streams if s.type == "audio")
|
||||
resampler = av.AudioResampler(format="fltp", layout="stereo", rate=target_sr)
|
||||
chunks: list = []
|
||||
for frame in container.decode(audio_stream):
|
||||
for resampled in resampler.resample(frame):
|
||||
chunks.append(resampled.to_ndarray())
|
||||
for resampled in resampler.resample(None):
|
||||
chunks.append(resampled.to_ndarray())
|
||||
container.close()
|
||||
if not chunks:
|
||||
raise RuntimeError(f"No audio frames decoded from {path}")
|
||||
waveform = np.concatenate(chunks, axis=-1)
|
||||
if waveform.ndim == 1:
|
||||
waveform = waveform[None, :]
|
||||
return torch.from_numpy(waveform).float()
|
||||
@@ -26,7 +26,7 @@ TURBO_T2V_1_3B = InferencePreset(
|
||||
"fps": 16,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
"negative_prompt": "",
|
||||
"negative_prompt": None,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -44,7 +44,7 @@ TURBO_T2V_14B = InferencePreset(
|
||||
"fps": 16,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
"negative_prompt": "",
|
||||
"negative_prompt": None,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -62,7 +62,7 @@ TURBO_I2V_A14B = InferencePreset(
|
||||
"fps": 16,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
"negative_prompt": "",
|
||||
"negative_prompt": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -17,8 +17,6 @@ import torch
|
||||
if TYPE_CHECKING:
|
||||
from torchcodec.decoders import VideoDecoder
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
@@ -192,21 +190,6 @@ class ForwardBatch:
|
||||
ltx2_stg_blocks_video: list[int] = field(default_factory=list)
|
||||
ltx2_stg_blocks_audio: list[int] = field(default_factory=list)
|
||||
|
||||
# Stable Audio (T2A): clip start/end in seconds. Parallels the
|
||||
# `SamplingParam` fields of the same name; the
|
||||
# `StableAudioConditioningStage` / `DecodingStage` read them.
|
||||
audio_start_in_s: float | None = None
|
||||
audio_end_in_s: float | None = None
|
||||
|
||||
# Stable Audio A2A variation + inpainting payloads (parallel to
|
||||
# `SamplingParam`). `Any` because we accept torch tensors or numpy
|
||||
# arrays the user supplies; the latent-prep stage normalises shapes.
|
||||
init_audio: Any = None
|
||||
init_audio_strength: float | None = None
|
||||
init_noise_level: float | None = None
|
||||
inpaint_audio: Any = None
|
||||
inpaint_mask: Any = None
|
||||
|
||||
n_tokens: int | None = None
|
||||
|
||||
# Other parameters that may be needed by specific schedulers
|
||||
@@ -223,9 +206,6 @@ class ForwardBatch:
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_decoded: list[torch.Tensor] | None = None
|
||||
|
||||
continuation_state: "ContinuationState | None" = None
|
||||
return_continuation_state: bool = False
|
||||
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@@ -1,193 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Preprocess Cosmos 2.5 overfit data into parquet format.
|
||||
|
||||
Encodes videos with the Cosmos (Wan-style) VAE and captions with the
|
||||
Reason1 (Qwen2.5-VL) text encoder into the t2v parquet schema.
|
||||
|
||||
Usage:
|
||||
CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
# --- Config ---
|
||||
NUM_FRAMES = 93 # 4*23+1 → 24 latent frames
|
||||
MAX_HEIGHT = 480
|
||||
MAX_WIDTH = 832
|
||||
TRAIN_FPS = 16.0
|
||||
|
||||
DATA_DIR = "data/cosmos_overfit"
|
||||
OUTPUT_DIR = "data/cosmos25_overfit_preprocessed"
|
||||
MODEL_REPO = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
# The VAE is architecturally identical to Cosmos Predict2;
|
||||
# use the Predict2 model for VAE since its weights are in
|
||||
# standard diffusers format.
|
||||
VAE_REPO = "nvidia/Cosmos-Predict2-2B-Video2World"
|
||||
|
||||
|
||||
def load_video(path: str, num_frames: int) -> torch.Tensor:
|
||||
"""Load video as [1, C, T, H, W] in [-1, 1]."""
|
||||
cap = cv2.VideoCapture(path)
|
||||
frames: list[np.ndarray] = []
|
||||
while len(frames) < num_frames:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
cap.release()
|
||||
|
||||
if len(frames) < num_frames:
|
||||
while len(frames) < num_frames:
|
||||
frames.append(frames[-1])
|
||||
|
||||
frames = frames[:num_frames]
|
||||
video = np.stack(frames, axis=0)
|
||||
video = torch.from_numpy(video).float()
|
||||
video = video / 127.5 - 1.0 # [0,255] -> [-1,1]
|
||||
video = video.permute(3, 0, 1, 2).unsqueeze(0) # [1,C,T,H,W]
|
||||
return video
|
||||
|
||||
|
||||
def main() -> None:
|
||||
device = torch.device("cuda:0")
|
||||
model_path = maybe_download_model(MODEL_REPO)
|
||||
|
||||
os.makedirs(OUTPUT_DIR, exist_ok=True)
|
||||
|
||||
# Load captions
|
||||
with open(os.path.join(DATA_DIR, "videos2caption.json")) as f:
|
||||
caption_data = json.load(f)
|
||||
|
||||
# --- Load VAE (Wan-style, same arch for Cosmos 2 and 2.5) ---
|
||||
print("Loading Cosmos VAE (AutoencoderKLWan)...")
|
||||
vae_path = maybe_download_model(VAE_REPO)
|
||||
from diffusers import AutoencoderKLWan
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
vae_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=torch.float16,
|
||||
).to(device).eval()
|
||||
print(f"VAE loaded "
|
||||
f"({sum(p.numel() for p in vae.parameters())/1e6:.0f}M)")
|
||||
|
||||
# --- Load Reason1 (Qwen2.5-VL) text encoder ---
|
||||
print("Loading Reason1 text encoder...")
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import (
|
||||
Cosmos25Config, )
|
||||
from fastvideo.models.encoders.reason1 import (
|
||||
Reason1TextEncoder, )
|
||||
pipeline_cfg = Cosmos25Config()
|
||||
text_enc_cfg = pipeline_cfg.text_encoder_configs[0]
|
||||
text_enc_path = os.path.join(model_path, "text_encoder")
|
||||
|
||||
# Instantiate Reason1TextEncoder with config and checkpoint
|
||||
text_encoder = Reason1TextEncoder(
|
||||
text_enc_cfg,
|
||||
checkpoint_path=text_enc_path,
|
||||
)
|
||||
# Load weights from safetensors into the meta-device model.
|
||||
# Materialize empty tensors in bf16 on the target device,
|
||||
# then overwrite with checkpoint weights.
|
||||
text_encoder = text_encoder.to_empty(device=device)
|
||||
text_encoder = text_encoder.to(torch.bfloat16)
|
||||
import glob
|
||||
from safetensors.torch import load_file
|
||||
sd: dict[str, torch.Tensor] = {}
|
||||
for sf in sorted(glob.glob(os.path.join(text_enc_path, "*.safetensors"))):
|
||||
sd.update(load_file(sf, device=str(device)))
|
||||
sd = {k: v.to(torch.bfloat16) for k, v in sd.items()}
|
||||
text_encoder.load_state_dict(sd, strict=False, assign=True)
|
||||
del sd
|
||||
torch.cuda.empty_cache()
|
||||
text_encoder = text_encoder.eval()
|
||||
print("Reason1 text encoder loaded")
|
||||
|
||||
# --- Process each video ---
|
||||
records = []
|
||||
for idx, item in enumerate(caption_data):
|
||||
video_name = item["path"]
|
||||
record_id = f"{idx:04d}_{video_name}"
|
||||
caption = item["cap"][0]
|
||||
video_path = os.path.join(DATA_DIR, "videos", video_name)
|
||||
|
||||
print(f"\nProcessing: {video_name}")
|
||||
print(f" Caption: {caption[:80]}...")
|
||||
|
||||
# Encode video
|
||||
video = load_video(video_path, NUM_FRAMES).to(device=device, dtype=torch.float16)
|
||||
print(f" Video shape: {video.shape}")
|
||||
|
||||
with torch.no_grad():
|
||||
latent_dist = vae.encode(video).latent_dist
|
||||
latent = latent_dist.mean.squeeze(0).float().cpu()
|
||||
print(f" Latent shape: {latent.shape}")
|
||||
|
||||
# Encode text with Reason1 (Qwen2.5-VL)
|
||||
with torch.no_grad():
|
||||
text_embedding = text_encoder.compute_text_embeddings(
|
||||
[caption],
|
||||
device=device,
|
||||
)
|
||||
text_embedding = text_embedding.squeeze(0).float().cpu()
|
||||
print(f" Text embedding shape: {text_embedding.shape}")
|
||||
|
||||
record = {
|
||||
"id": record_id,
|
||||
"vae_latent_bytes": latent.numpy().tobytes(),
|
||||
"vae_latent_shape": list(latent.shape),
|
||||
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
|
||||
"text_embedding_bytes": (text_embedding.numpy().tobytes()),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
|
||||
"file_name": video_name,
|
||||
"caption": caption,
|
||||
"media_type": "video",
|
||||
"width": MAX_WIDTH,
|
||||
"height": MAX_HEIGHT,
|
||||
"num_frames": NUM_FRAMES,
|
||||
"duration_sec": NUM_FRAMES / TRAIN_FPS,
|
||||
"fps": TRAIN_FPS,
|
||||
}
|
||||
records.append(record)
|
||||
|
||||
# Clean up
|
||||
del text_encoder, vae
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Write parquet
|
||||
table = pa.table(
|
||||
{k: [r[k] for r in records]
|
||||
for k in records[0]},
|
||||
schema=pyarrow_schema_t2v,
|
||||
)
|
||||
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
|
||||
pq.write_table(table, output_path)
|
||||
print(f"\nWrote {len(records)} records to {output_path}")
|
||||
|
||||
# Write T2W validation prompts (no image_path for T2W)
|
||||
val_prompts = {
|
||||
"data": [{
|
||||
"caption": item["cap"][0],
|
||||
} for item in caption_data],
|
||||
}
|
||||
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
|
||||
with open(val_path, "w") as f:
|
||||
json.dump(val_prompts, f, indent=2)
|
||||
print(f"Wrote validation prompts to {val_path}")
|
||||
|
||||
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,196 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Preprocess Cosmos-Predict2 overfit data into parquet format.
|
||||
|
||||
Encodes videos with the Cosmos (Wan-style) VAE and captions with the
|
||||
single T5 Large text encoder into the t2v parquet schema expected by
|
||||
the training framework.
|
||||
|
||||
Usage:
|
||||
CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos_overfit.py
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.encoders import T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.base import BaseEncoderOutput
|
||||
from fastvideo.configs.pipelines.cosmos import t5_large_postprocess_text
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
# --- Config ---
|
||||
NUM_FRAMES = 93 # 4*23+1 for temporal compression ratio 4 -> 24 latent frames
|
||||
MAX_HEIGHT = 480
|
||||
MAX_WIDTH = 832
|
||||
TRAIN_FPS = 16.0
|
||||
|
||||
DATA_DIR = "data/cosmos_overfit"
|
||||
OUTPUT_DIR = "data/cosmos_overfit_preprocessed"
|
||||
MODEL_REPO = "nvidia/Cosmos-Predict2-2B-Video2World"
|
||||
|
||||
|
||||
def load_video(path: str, num_frames: int) -> torch.Tensor:
|
||||
"""Load video as [1, C, T, H, W] in [-1, 1]."""
|
||||
cap = cv2.VideoCapture(path)
|
||||
frames: list[np.ndarray] = []
|
||||
while len(frames) < num_frames:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
cap.release()
|
||||
|
||||
if len(frames) < num_frames:
|
||||
# Repeat last frame to fill
|
||||
while len(frames) < num_frames:
|
||||
frames.append(frames[-1])
|
||||
|
||||
frames = frames[:num_frames]
|
||||
video = np.stack(frames, axis=0)
|
||||
video = torch.from_numpy(video).float()
|
||||
video = video / 127.5 - 1.0 # [0,255] -> [-1,1]
|
||||
video = video.permute(3, 0, 1, 2).unsqueeze(0) # [1,C,T,H,W]
|
||||
return video
|
||||
|
||||
|
||||
def main() -> None:
|
||||
device = torch.device("cuda:0")
|
||||
model_path = maybe_download_model(MODEL_REPO)
|
||||
|
||||
os.makedirs(OUTPUT_DIR, exist_ok=True)
|
||||
|
||||
# Load captions
|
||||
with open(os.path.join(DATA_DIR, "videos2caption.json")) as f:
|
||||
caption_data = json.load(f)
|
||||
|
||||
# --- Load VAE ---
|
||||
# Cosmos-Predict2-2B ships a Wan-style VAE; vae/config.json declares
|
||||
# `_class_name: AutoencoderKLWan`, so diffusers can load it directly.
|
||||
print("Loading Cosmos VAE (AutoencoderKLWan)...")
|
||||
from diffusers import AutoencoderKLWan
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
model_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=torch.float16,
|
||||
).to(device).eval()
|
||||
print(f"VAE loaded ({sum(p.numel() for p in vae.parameters())/1e6:.0f}M)")
|
||||
|
||||
# --- Load T5 Large text encoder ---
|
||||
print("Loading T5 Large text encoder...")
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
|
||||
t5_cfg = T5LargeConfig()
|
||||
tok_kwargs = dict(t5_cfg.tokenizer_kwargs)
|
||||
tokenizer = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer"))
|
||||
text_encoder = T5EncoderModel.from_pretrained(
|
||||
os.path.join(model_path, "text_encoder"),
|
||||
torch_dtype=torch.bfloat16,
|
||||
).to(device).eval()
|
||||
|
||||
# --- Process each video ---
|
||||
records = []
|
||||
for idx, item in enumerate(caption_data):
|
||||
video_name = item["path"]
|
||||
record_id = f"{idx:04d}_{video_name}"
|
||||
caption = item["cap"][0]
|
||||
video_path = os.path.join(DATA_DIR, "videos", video_name)
|
||||
|
||||
print(f"\nProcessing: {video_name}")
|
||||
print(f" Caption: {caption[:80]}...")
|
||||
|
||||
# Encode video
|
||||
video = load_video(video_path, NUM_FRAMES).to(device=device, dtype=torch.float16)
|
||||
print(f" Video shape: {video.shape}")
|
||||
|
||||
with torch.no_grad():
|
||||
latent_dist = vae.encode(video).latent_dist
|
||||
# Cast to fp32 — dataloader hardcodes np.float32
|
||||
latent = latent_dist.mean.squeeze(0).float().cpu()
|
||||
print(f" Latent shape: {latent.shape}")
|
||||
|
||||
# Encode text with T5
|
||||
with torch.no_grad():
|
||||
inputs = tokenizer(caption, **tok_kwargs).to(device)
|
||||
outputs = text_encoder(**inputs)
|
||||
enc_out = BaseEncoderOutput(
|
||||
last_hidden_state=outputs.last_hidden_state,
|
||||
attention_mask=inputs["attention_mask"],
|
||||
)
|
||||
# [1, max_len, 1024], zeros beyond real length
|
||||
t5_embed = t5_large_postprocess_text(enc_out).squeeze(0)
|
||||
# Trim to real sequence length so dataloader's pad() builds
|
||||
# the correct attention mask.
|
||||
real_len = int(inputs["attention_mask"].sum().item())
|
||||
text_embedding = t5_embed[:real_len].float().cpu() # [seq, 1024]
|
||||
print(f" Text embedding shape: {text_embedding.shape}")
|
||||
|
||||
record = {
|
||||
"id": record_id,
|
||||
"vae_latent_bytes": latent.numpy().tobytes(),
|
||||
"vae_latent_shape": list(latent.shape),
|
||||
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
|
||||
"text_embedding_bytes": text_embedding.numpy().tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
|
||||
"file_name": video_name,
|
||||
"caption": caption,
|
||||
"media_type": "video",
|
||||
"width": MAX_WIDTH,
|
||||
"height": MAX_HEIGHT,
|
||||
"num_frames": NUM_FRAMES,
|
||||
"duration_sec": NUM_FRAMES / TRAIN_FPS,
|
||||
"fps": TRAIN_FPS,
|
||||
}
|
||||
records.append(record)
|
||||
|
||||
# Clean up encoders
|
||||
del text_encoder, tokenizer, vae
|
||||
|
||||
# Write parquet
|
||||
table = pa.table(
|
||||
{k: [r[k] for r in records]
|
||||
for k in records[0]},
|
||||
schema=pyarrow_schema_t2v,
|
||||
)
|
||||
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
|
||||
pq.write_table(table, output_path)
|
||||
print(f"\nWrote {len(records)} records to {output_path}")
|
||||
|
||||
# Extract first frame from first video as V2W conditioning image
|
||||
import cv2
|
||||
first_video = os.path.join(DATA_DIR, "videos", caption_data[0]["path"])
|
||||
cap = cv2.VideoCapture(first_video)
|
||||
ret, frame = cap.read()
|
||||
cap.release()
|
||||
cond_frame_path = os.path.join(OUTPUT_DIR, "cond_frame.png")
|
||||
if ret:
|
||||
cv2.imwrite(cond_frame_path, frame)
|
||||
print(f"Saved conditioning frame to {cond_frame_path}")
|
||||
|
||||
# Write validation prompts for callback
|
||||
# Wrap in "data" key — ValidationDataset expects field="data"
|
||||
# Use "caption" field — ValidationDataset aliases it to "prompt"
|
||||
# Include image_path for V2W conditioning during validation
|
||||
val_prompts = {
|
||||
"data": [{
|
||||
"caption": item["cap"][0],
|
||||
"image_path": "cond_frame.png",
|
||||
} for item in caption_data]
|
||||
}
|
||||
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
|
||||
with open(val_path, "w") as f:
|
||||
json.dump(val_prompts, f, indent=2)
|
||||
print(f"Wrote validation prompts to {val_path}")
|
||||
|
||||
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -518,42 +518,13 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
|
||||
class CosmosDenoisingStage(DenoisingStage):
|
||||
"""Denoising stage for Cosmos models.
|
||||
|
||||
Uses FlowMatchEulerDiscreteScheduler with manual EDM
|
||||
preconditioning (c_in, c_skip, c_out) to match the
|
||||
pretrained Cosmos model's training convention.
|
||||
"""
|
||||
Denoising stage for Cosmos models using FlowMatchEulerDiscreteScheduler.
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None) -> None:
|
||||
super().__init__(transformer, scheduler, pipeline)
|
||||
|
||||
def _run_transformer(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
condition_mask: torch.Tensor,
|
||||
padding_mask: torch.Tensor,
|
||||
target_dtype: torch.dtype,
|
||||
step_index: int,
|
||||
batch: ForwardBatch,
|
||||
) -> torch.Tensor:
|
||||
with set_forward_context(
|
||||
current_timestep=step_index,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
return self.transformer(
|
||||
hidden_states=hidden_states.to(target_dtype),
|
||||
timestep=timestep.to(target_dtype),
|
||||
encoder_hidden_states=encoder_hidden_states.to(target_dtype),
|
||||
fps=24,
|
||||
condition_mask=condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
@@ -562,188 +533,199 @@ class CosmosDenoisingStage(DenoisingStage):
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"],
|
||||
fastvideo_args,
|
||||
)
|
||||
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
if hasattr(self.transformer, "module"):
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
{
|
||||
"generator": batch.generator,
|
||||
"eta": batch.eta
|
||||
},
|
||||
)
|
||||
|
||||
if hasattr(self.transformer, 'module'):
|
||||
transformer_dtype = next(self.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
target_dtype = transformer_dtype
|
||||
autocast_enabled = (target_dtype != torch.float32 and not fastvideo_args.disable_autocast)
|
||||
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_cfg = (batch.do_classifier_free_guidance and batch.negative_prompt_embeds is not None)
|
||||
|
||||
sigma_data = float(getattr(self.scheduler.config, "sigma_data", 1.0))
|
||||
sigma_max = 80.0
|
||||
sigma_min = 0.002
|
||||
sigma_data = 1.0
|
||||
final_sigmas_type = "sigma_min"
|
||||
|
||||
self.scheduler.set_timesteps(
|
||||
num_inference_steps,
|
||||
device=latents.device,
|
||||
)
|
||||
if self.scheduler is not None:
|
||||
self.scheduler.register_to_config(
|
||||
sigma_max=sigma_max,
|
||||
sigma_min=sigma_min,
|
||||
sigma_data=sigma_data,
|
||||
final_sigmas_type=final_sigmas_type,
|
||||
)
|
||||
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=latents.device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# Clamp terminal sigma to sigma_min (avoid zero).
|
||||
if (hasattr(self.scheduler.config, "final_sigmas_type")
|
||||
if (hasattr(self.scheduler.config, 'final_sigmas_type')
|
||||
and self.scheduler.config.final_sigmas_type == "sigma_min" and len(self.scheduler.sigmas) > 1):
|
||||
self.scheduler.sigmas[-1] = self.scheduler.sigmas[-2]
|
||||
|
||||
conditioning_latents = getattr(
|
||||
batch,
|
||||
"conditioning_latents",
|
||||
None,
|
||||
)
|
||||
cond_indicator = getattr(batch, "cond_indicator", None)
|
||||
uncond_indicator = getattr(
|
||||
batch,
|
||||
"uncond_indicator",
|
||||
None,
|
||||
)
|
||||
conditioning_latents = getattr(batch, 'conditioning_latents', None)
|
||||
unconditioning_latents = conditioning_latents
|
||||
|
||||
augment_sigma = torch.tensor(
|
||||
[0.001],
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
padding_mask = torch.zeros(
|
||||
1,
|
||||
1,
|
||||
batch.height,
|
||||
batch.width,
|
||||
device=latents.device,
|
||||
dtype=target_dtype,
|
||||
)
|
||||
|
||||
condition_mask = (batch.cond_mask.to(target_dtype)
|
||||
if hasattr(batch, "cond_mask") and batch.cond_mask is not None else None)
|
||||
uncond_condition_mask = (batch.uncond_mask.to(target_dtype)
|
||||
if hasattr(batch, "uncond_mask") and batch.uncond_mask is not None else condition_mask)
|
||||
if condition_mask is None:
|
||||
b, c, tf, h, w = latents.shape
|
||||
condition_mask = torch.zeros(
|
||||
b,
|
||||
1,
|
||||
tf,
|
||||
h,
|
||||
w,
|
||||
device=latents.device,
|
||||
dtype=target_dtype,
|
||||
)
|
||||
uncond_condition_mask = condition_mask
|
||||
|
||||
with self.progress_bar(total=num_inference_steps, ) as progress_bar:
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if hasattr(self, "interrupt") and self.interrupt:
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
|
||||
sigma = self.scheduler.sigmas[i]
|
||||
is_aug_greater = bool(augment_sigma >= sigma)
|
||||
current_sigma = self.scheduler.sigmas[i]
|
||||
current_t = current_sigma / (current_sigma + 1)
|
||||
c_in = 1 - current_t
|
||||
c_skip = 1 - current_t
|
||||
c_out = -current_t
|
||||
|
||||
# EDM preconditioning coefficients.
|
||||
c_in = 1.0 / (sigma**2 + sigma_data**2)**0.5
|
||||
c_in_aug = 1.0 / (augment_sigma**2 + sigma_data**2)**0.5
|
||||
c_skip = sigma_data**2 / (sigma**2 + sigma_data**2)
|
||||
c_out = (sigma * sigma_data / (sigma**2 + sigma_data**2)**0.5)
|
||||
timestep = current_t.view(1, 1, 1, 1, 1).expand(latents.size(0), -1, latents.size(2), -1,
|
||||
-1) # [B, 1, T, 1, 1]
|
||||
|
||||
# The model expects timestep = sigma * 1000
|
||||
# (FlowMatchEulerDiscreteScheduler convention).
|
||||
timestep_expanded = t.expand(latents.shape[0], ).to(target_dtype)
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
|
||||
with torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled,
|
||||
):
|
||||
# --- Conditioning frame injection ---
|
||||
cur_ci = (cond_indicator * 0 if cond_indicator is not None and is_aug_greater else cond_indicator)
|
||||
cond_latent = latents * c_in
|
||||
|
||||
cond_latent = latents.clone()
|
||||
if (cur_ci is not None and conditioning_latents is not None):
|
||||
cn = torch.randn_like(
|
||||
latents,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
cf = (conditioning_latents + cn * augment_sigma[:, None, None, None, None])
|
||||
cf = cf * c_in_aug / c_in
|
||||
cond_latent = (cur_ci * cf + (1 - cur_ci) * cond_latent)
|
||||
|
||||
# Manual EDM input scaling.
|
||||
model_input = cond_latent * c_in
|
||||
|
||||
noise_pred_cond = self._run_transformer(
|
||||
model_input,
|
||||
timestep_expanded,
|
||||
batch.prompt_embeds[0],
|
||||
condition_mask,
|
||||
padding_mask,
|
||||
target_dtype,
|
||||
i,
|
||||
batch,
|
||||
)
|
||||
|
||||
# EDM output → x0 prediction.
|
||||
cond_x0 = (c_skip * latents + c_out * noise_pred_cond.float())
|
||||
if (cur_ci is not None and conditioning_latents is not None):
|
||||
cond_x0 = (cur_ci * conditioning_latents + (1 - cur_ci) * cond_x0)
|
||||
|
||||
# --- CFG: unconditional pass ---
|
||||
if do_cfg:
|
||||
cur_ui = (uncond_indicator *
|
||||
0 if uncond_indicator is not None and is_aug_greater else uncond_indicator)
|
||||
|
||||
uncond_latent = latents.clone()
|
||||
if (cur_ui is not None and conditioning_latents is not None):
|
||||
un = torch.randn_like(
|
||||
latents,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
uf = (conditioning_latents + un * augment_sigma[:, None, None, None, None])
|
||||
uf = uf * c_in_aug / c_in
|
||||
uncond_latent = (cur_ui * uf + (1 - cur_ui) * uncond_latent)
|
||||
|
||||
uncond_input = uncond_latent * c_in
|
||||
|
||||
noise_pred_uncond = (self._run_transformer(
|
||||
uncond_input,
|
||||
timestep_expanded,
|
||||
batch.negative_prompt_embeds[0],
|
||||
uncond_condition_mask,
|
||||
padding_mask,
|
||||
target_dtype,
|
||||
i,
|
||||
if hasattr(
|
||||
batch,
|
||||
))
|
||||
|
||||
uncond_x0 = (c_skip * latents + c_out * noise_pred_uncond.float())
|
||||
if (cur_ui is not None and conditioning_latents is not None):
|
||||
uncond_x0 = (cur_ui * conditioning_latents + (1 - cur_ui) * uncond_x0)
|
||||
|
||||
final_x0 = (cond_x0 + guidance_scale * (cond_x0 - uncond_x0))
|
||||
'cond_indicator') and batch.cond_indicator is not None and conditioning_latents is not None:
|
||||
cond_latent = batch.cond_indicator * conditioning_latents + (1 -
|
||||
batch.cond_indicator) * cond_latent
|
||||
else:
|
||||
final_x0 = cond_x0
|
||||
logger.warning(
|
||||
"Step %s: Missing conditioning data - cond_indicator: %s, conditioning_latents: %s", i,
|
||||
hasattr(batch, 'cond_indicator'), conditioning_latents is not None)
|
||||
|
||||
# Convert x0 to velocity for
|
||||
# FlowMatchEulerDiscreteScheduler.
|
||||
velocity = (latents - final_x0) / sigma.clamp(min=1e-6)
|
||||
cond_latent = cond_latent.to(target_dtype)
|
||||
|
||||
latents = self.scheduler.step(
|
||||
velocity,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
cond_timestep = timestep
|
||||
if hasattr(batch, 'cond_indicator') and batch.cond_indicator is not None:
|
||||
sigma_conditioning = 0.0001
|
||||
t_conditioning = sigma_conditioning / (sigma_conditioning + 1)
|
||||
cond_timestep = batch.cond_indicator * t_conditioning + (1 - batch.cond_indicator) * timestep
|
||||
cond_timestep = cond_timestep.to(target_dtype)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
# Use conditioning masks from CosmosLatentPreparationStage
|
||||
condition_mask = batch.cond_mask.to(target_dtype) if hasattr(batch, 'cond_mask') else None
|
||||
padding_mask = torch.zeros(1,
|
||||
1,
|
||||
batch.height,
|
||||
batch.width,
|
||||
device=cond_latent.device,
|
||||
dtype=target_dtype)
|
||||
|
||||
# Fallback if masks not available
|
||||
if condition_mask is None:
|
||||
batch_size, num_channels, num_frames, height, width = cond_latent.shape
|
||||
condition_mask = torch.zeros(batch_size,
|
||||
1,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
device=cond_latent.device,
|
||||
dtype=target_dtype)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=cond_latent,
|
||||
timestep=cond_timestep.to(target_dtype),
|
||||
encoder_hidden_states=batch.prompt_embeds[0].to(target_dtype),
|
||||
fps=24, # TODO: get fps from batch or config
|
||||
condition_mask=condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
cond_pred = (c_skip * latents + c_out * noise_pred.float()).to(target_dtype)
|
||||
|
||||
if hasattr(
|
||||
batch,
|
||||
'cond_indicator') and batch.cond_indicator is not None and conditioning_latents is not None:
|
||||
cond_pred = batch.cond_indicator * conditioning_latents + (1 - batch.cond_indicator) * cond_pred
|
||||
|
||||
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds is not None:
|
||||
uncond_latent = latents * c_in
|
||||
|
||||
if hasattr(batch, 'uncond_indicator'
|
||||
) and batch.uncond_indicator is not None and unconditioning_latents is not None:
|
||||
uncond_latent = batch.uncond_indicator * unconditioning_latents + (
|
||||
1 - batch.uncond_indicator) * uncond_latent
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
uncond_condition_mask = batch.uncond_mask.to(target_dtype) if hasattr(
|
||||
batch, 'uncond_mask') and batch.uncond_mask is not None else condition_mask
|
||||
|
||||
uncond_timestep = timestep
|
||||
if hasattr(batch, 'uncond_indicator') and batch.uncond_indicator is not None:
|
||||
sigma_conditioning = 0.0001
|
||||
t_conditioning = sigma_conditioning / (sigma_conditioning + 1)
|
||||
uncond_timestep = batch.uncond_indicator * t_conditioning + (
|
||||
1 - batch.uncond_indicator) * timestep
|
||||
uncond_timestep = uncond_timestep.to(target_dtype)
|
||||
|
||||
noise_pred_uncond = self.transformer(
|
||||
hidden_states=uncond_latent.to(target_dtype),
|
||||
timestep=uncond_timestep.to(target_dtype),
|
||||
encoder_hidden_states=batch.negative_prompt_embeds[0].to(target_dtype),
|
||||
fps=24, # TODO: get fps from batch or config
|
||||
condition_mask=uncond_condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
uncond_pred = (c_skip * latents + c_out * noise_pred_uncond.float()).to(target_dtype)
|
||||
|
||||
if hasattr(batch, 'uncond_indicator'
|
||||
) and batch.uncond_indicator is not None and unconditioning_latents is not None:
|
||||
uncond_pred = batch.uncond_indicator * unconditioning_latents + (
|
||||
1 - batch.uncond_indicator) * uncond_pred
|
||||
|
||||
guidance_diff = cond_pred - uncond_pred
|
||||
final_pred = cond_pred + guidance_scale * guidance_diff
|
||||
else:
|
||||
final_pred = cond_pred
|
||||
|
||||
# Convert to noise for scheduler step
|
||||
if current_sigma > 1e-8:
|
||||
noise_for_scheduler = (latents - final_pred) / current_sigma
|
||||
else:
|
||||
logger.warning("Step %s: current_sigma too small (%s), using final_pred directly", i, current_sigma)
|
||||
noise_for_scheduler = final_pred
|
||||
|
||||
if torch.isnan(noise_for_scheduler).sum() > 0:
|
||||
logger.error("Step %s: NaN detected in noise_for_scheduler, sum: %s", i,
|
||||
noise_for_scheduler.float().sum().item())
|
||||
logger.error("Step %s: latents sum: %s, final_pred sum: %s, current_sigma: %s", i,
|
||||
latents.float().sum().item(),
|
||||
final_pred.float().sum().item(), current_sigma)
|
||||
|
||||
latents = self.scheduler.step(noise_for_scheduler, t, latents, **extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
batch.latents = latents
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
@@ -789,21 +771,16 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
|
||||
},
|
||||
)
|
||||
|
||||
# Detect the actual weight dtype. FSDP-wrapped models may
|
||||
# report fp32 via next(parameters()) even when the physical
|
||||
# weights are bf16. Walk through parameters to find one
|
||||
# that is NOT fp32 (the real checkpoint dtype).
|
||||
target_dtype = torch.bfloat16 # safe default for Cosmos 2.5
|
||||
for p in self.transformer.parameters():
|
||||
if p.dtype != torch.float32:
|
||||
target_dtype = p.dtype
|
||||
break
|
||||
if hasattr(self.transformer, 'module'):
|
||||
transformer_dtype = next(self.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
target_dtype = transformer_dtype
|
||||
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents = batch.latents
|
||||
if latents is None:
|
||||
raise ValueError("latents must be provided for "
|
||||
"Cosmos25DenoisingStage")
|
||||
raise ValueError("latents must be provided for Cosmos25DenoisingStage")
|
||||
guidance_scale = batch.guidance_scale
|
||||
|
||||
if batch.timesteps is None:
|
||||
|
||||
@@ -524,22 +524,6 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
image = resize(image, height, width, resize_mode=resize_mode)
|
||||
image = pil_to_numpy(image) # to np
|
||||
image = numpy_to_pt(image) # to pt
|
||||
elif isinstance(image, torch.Tensor):
|
||||
# VideoTransformStage delivers uint8 [0, 255] frames via batch.pil_image
|
||||
# for the I2V preprocessing path. Convert here (not at the source) because
|
||||
# batch.pil_image is also consumed as uint8 by ImageEncodingStage (HF
|
||||
# processor does its own rescale) and by record_schema.py for parquet.
|
||||
if image.dtype == torch.uint8:
|
||||
image = image.float() / 255.0
|
||||
elif not image.dtype.is_floating_point:
|
||||
raise ValueError(f"preprocess() expected uint8 or float tensor, got {image.dtype}")
|
||||
image_min = image.min()
|
||||
image_max = image.max()
|
||||
if image_max > 1.0 + 1e-4 or image_min < -1.0 - 1e-4:
|
||||
raise ValueError("preprocess() expected tensor in [0, 1] or [-1, 1], got "
|
||||
f"range [{image_min.item():.3f}, {image_max.item():.3f}]")
|
||||
else:
|
||||
raise TypeError(f"preprocess() expected PIL.Image or torch.Tensor, got {type(image)}")
|
||||
|
||||
do_normalize = True
|
||||
if image.min() < 0:
|
||||
|
||||
@@ -48,7 +48,6 @@ from fastvideo.configs.pipelines.wan import (
|
||||
WanT2V720PConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.sd35 import SD35Config
|
||||
from fastvideo.configs.pipelines.stable_audio import (StableAudioOpenSmallConfig, StableAudioT2AConfig)
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
@@ -243,47 +242,6 @@ def _register_configs() -> None:
|
||||
default_preset="ltx2_distilled",
|
||||
)
|
||||
|
||||
# Stable Audio Open (text-to-audio). Both variants must be loaded
|
||||
# from the FastVideo-curated converted Diffusers-format repos —
|
||||
# the upstream `stabilityai/stable-audio-open-{1.0,small}` repos
|
||||
# ship `model.safetensors` as a single monolithic checkpoint with
|
||||
# no per-component subfolders our standard loader can consume. See
|
||||
# `scripts/checkpoint_conversion/stable_audio_to_diffusers.py`.
|
||||
# NOTE: WorkloadType has no T2A variant yet (REVIEW item 28); using
|
||||
# T2V as the placeholder until the enum is extended.
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=StableAudioT2AConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
],
|
||||
# Substring match against HF cache snapshot paths (the lookup
|
||||
# runs on the resolved local directory, which uses `--` between
|
||||
# org and repo: `models--FastVideo--stable-audio-open-1.0-Diffusers`).
|
||||
model_detectors=[
|
||||
lambda path: "stable-audio-open-1" in path.lower(),
|
||||
],
|
||||
model_family="stable_audio",
|
||||
default_preset="stable_audio_open_1_0_base",
|
||||
)
|
||||
# Small variant uses its own `pipeline_config_cls` so it picks up
|
||||
# the smaller (524288-sample) training window in `sample_size` /
|
||||
# `max_audio_duration_s`.
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=StableAudioOpenSmallConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/stable-audio-open-small-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "stable-audio-open-small" in path.lower(),
|
||||
],
|
||||
model_family="stable_audio",
|
||||
default_preset="stable_audio_open_small",
|
||||
)
|
||||
|
||||
# Hunyuan 1.5 (specific)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
@@ -828,8 +786,6 @@ def _register_presets() -> None:
|
||||
ALL_PRESETS as MATRIXGAME_PRESETS, )
|
||||
from fastvideo.pipelines.basic.sd35.presets import (
|
||||
ALL_PRESETS as SD35_PRESETS, )
|
||||
from fastvideo.pipelines.basic.stable_audio.presets import (
|
||||
ALL_PRESETS as STABLE_AUDIO_PRESETS, )
|
||||
from fastvideo.pipelines.basic.turbodiffusion.presets import (
|
||||
ALL_PRESETS as TURBODIFFUSION_PRESETS, )
|
||||
from fastvideo.pipelines.basic.wan.presets import (
|
||||
@@ -847,7 +803,6 @@ def _register_presets() -> None:
|
||||
LTX2_PRESETS,
|
||||
MATRIXGAME_PRESETS,
|
||||
SD35_PRESETS,
|
||||
STABLE_AUDIO_PRESETS,
|
||||
TURBODIFFUSION_PRESETS,
|
||||
WAN_PRESETS,
|
||||
)
|
||||
|
||||
@@ -519,7 +519,7 @@ def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch)
|
||||
cli_main.main()
|
||||
|
||||
|
||||
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, monkeypatch):
|
||||
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
|
||||
config_path = tmp_path / "serve-streaming.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
@@ -530,21 +530,9 @@ def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, mo
|
||||
)
|
||||
args, _ = _parse_serve_args(["--config", str(config_path)])
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_run_server(serve_config, *, generator=None):
|
||||
captured["serve_config"] = serve_config
|
||||
|
||||
def fail_if_called(*_args, **_kwargs):
|
||||
raise AssertionError("OpenAI server must not run when streaming is set")
|
||||
|
||||
monkeypatch.setattr(streaming_server, "run_server", fake_run_server)
|
||||
monkeypatch.setattr(api_server, "run_server", fail_if_called)
|
||||
ServeSubcommand().cmd(args)
|
||||
|
||||
serve_config = captured["serve_config"]
|
||||
assert serve_config.streaming is not None
|
||||
assert serve_config.streaming.stream_mode == "av_fmp4"
|
||||
with pytest.raises(NotImplementedError,
|
||||
match="streaming server is not implemented"):
|
||||
ServeSubcommand().cmd(args)
|
||||
|
||||
|
||||
def test_streaming_run_server_rejects_missing_streaming_block():
|
||||
|
||||
@@ -155,51 +155,6 @@ class TestLegacyLtx2VaeTilingTranslation:
|
||||
assert "ltx2_vae_tiling" not in args.kwargs
|
||||
|
||||
|
||||
class TestLegacyTextEncoderCompileTranslation:
|
||||
"""``enable_torch_compile_text_encoder`` flat kwarg promotes to
|
||||
``generator.engine.compile.text_encoder_enabled``; reverse direction
|
||||
emits the legacy name back onto the FastVideoArgs kwargs dict so
|
||||
realtime-runtime consumers can read it before FastVideoArgs filters
|
||||
unknown fields."""
|
||||
|
||||
def test_forward_routes_to_compile_text_encoder_enabled(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"enable_torch_compile_text_encoder": True},
|
||||
)
|
||||
assert config.engine.compile.text_encoder_enabled is True
|
||||
|
||||
def test_false_round_trips(self) -> None:
|
||||
config = legacy_from_pretrained_to_config(
|
||||
"/models/ltx2",
|
||||
{"enable_torch_compile_text_encoder": False},
|
||||
)
|
||||
assert config.engine.compile.text_encoder_enabled is False
|
||||
|
||||
def test_unset_stays_none(self) -> None:
|
||||
config = legacy_from_pretrained_to_config("/models/ltx2", {})
|
||||
assert config.engine.compile.text_encoder_enabled is None
|
||||
|
||||
def test_reverse_emits_legacy_name(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(
|
||||
CompileConfig(text_encoder_enabled=True)),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert args.kwargs["enable_torch_compile_text_encoder"] is True
|
||||
|
||||
def test_reverse_unset_skips_key(self, monkeypatch) -> None:
|
||||
_stub_fastvideo_args_from_kwargs(monkeypatch)
|
||||
config = GeneratorConfig(
|
||||
model_path="/models/ltx2",
|
||||
engine=_engine_with_compile(CompileConfig()),
|
||||
)
|
||||
args = generator_config_to_fastvideo_args(config)
|
||||
assert "enable_torch_compile_text_encoder" not in args.kwargs
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Helpers
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
@@ -1,250 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the typed LTX-2 continuation state.
|
||||
|
||||
Covers:
|
||||
|
||||
* round-trip through :class:`ContinuationState` (inline and blob-backed)
|
||||
* payload is JSON-serializable (Dynamo RPC / HTTP client constraint)
|
||||
* kind / schema_version validation on deserialization
|
||||
* compat-layer validation (known kinds, payload shape)
|
||||
* round-trip through :func:`request_to_sampling_param` attaches the
|
||||
state to the resulting :class:`SamplingParam` without losing fidelity
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# Importing compat first, then the LTX-2 module, exercises the
|
||||
# self-registration side effect on import (important for the API
|
||||
# test suite where the pipeline package isn't otherwise imported).
|
||||
from fastvideo.api import compat as api_compat # noqa: F401
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
OutputConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import InMemoryBlobStore
|
||||
from fastvideo.pipelines.basic.ltx2.continuation import (
|
||||
LTX2_CONTINUATION_KIND,
|
||||
LTX2_CONTINUATION_SCHEMA_VERSION,
|
||||
LTX2ContinuationState,
|
||||
)
|
||||
|
||||
|
||||
def _make_typed_state() -> LTX2ContinuationState:
|
||||
return LTX2ContinuationState(
|
||||
segment_index=3,
|
||||
video_frames=[
|
||||
(np.ones((64, 64, 3), dtype=np.uint8) * (i * 10)) for i in range(4)
|
||||
],
|
||||
video_conditioning_frame_idx=9,
|
||||
video_conditioning_strength=0.75,
|
||||
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.float32),
|
||||
audio_sample_rate=24000,
|
||||
audio_conditioning_num_frames=5,
|
||||
audio_conditioning_strength=0.5,
|
||||
video_position_offset_sec=0.125,
|
||||
metadata={"note": "unit-test"},
|
||||
)
|
||||
|
||||
|
||||
class TestRoundTrip:
|
||||
"""Round-trip through :class:`ContinuationState` preserves all fields."""
|
||||
|
||||
def test_kind_and_schema_version(self):
|
||||
state = _make_typed_state().to_continuation_state()
|
||||
assert state.kind == LTX2_CONTINUATION_KIND
|
||||
assert state.payload["schema_version"] == LTX2_CONTINUATION_SCHEMA_VERSION
|
||||
|
||||
def test_inline_roundtrip_preserves_scalars(self):
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.segment_index == original.segment_index
|
||||
assert restored.video_conditioning_frame_idx == (
|
||||
original.video_conditioning_frame_idx)
|
||||
assert restored.video_conditioning_strength == (
|
||||
original.video_conditioning_strength)
|
||||
assert restored.audio_sample_rate == original.audio_sample_rate
|
||||
assert restored.audio_conditioning_num_frames == (
|
||||
original.audio_conditioning_num_frames)
|
||||
assert restored.audio_conditioning_strength == (
|
||||
original.audio_conditioning_strength)
|
||||
assert restored.video_position_offset_sec == (
|
||||
original.video_position_offset_sec)
|
||||
assert restored.metadata == original.metadata
|
||||
|
||||
def test_inline_roundtrip_preserves_video_frames(self):
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.video_frames is not None
|
||||
assert len(restored.video_frames) == len(original.video_frames)
|
||||
for before, after in zip(original.video_frames,
|
||||
restored.video_frames):
|
||||
np.testing.assert_array_equal(before, after)
|
||||
|
||||
def test_inline_roundtrip_preserves_audio_latents(self):
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.audio_latents is not None
|
||||
assert tuple(restored.audio_latents.shape) == tuple(
|
||||
original.audio_latents.shape)
|
||||
assert restored.audio_latents.dtype == original.audio_latents.dtype
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, original.audio_latents)
|
||||
|
||||
def test_payload_is_json_serializable(self):
|
||||
envelope = _make_typed_state().to_continuation_state()
|
||||
# json.dumps must not raise — required for Dynamo RPC transport
|
||||
# and HTTP client round-trip.
|
||||
reserialized = json.loads(json.dumps(envelope.payload))
|
||||
restored = LTX2ContinuationState.from_continuation_state(
|
||||
ContinuationState(
|
||||
kind=envelope.kind,
|
||||
payload=reserialized,
|
||||
))
|
||||
assert restored.segment_index == 3
|
||||
|
||||
def test_bf16_audio_latents_preserved(self):
|
||||
"""safetensors serialization must preserve bf16 dtype (numpy
|
||||
has no bf16, so a raw-bytes path would silently promote)."""
|
||||
state = LTX2ContinuationState(
|
||||
segment_index=0,
|
||||
audio_latents=torch.randn(1, 4, 16, 64, dtype=torch.bfloat16),
|
||||
)
|
||||
envelope = state.to_continuation_state()
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.audio_latents is not None
|
||||
assert restored.audio_latents.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, state.audio_latents)
|
||||
|
||||
|
||||
class TestBlobIndirection:
|
||||
"""Large tensors live in the :class:`BlobStore` instead of the payload."""
|
||||
|
||||
def test_threshold_triggers_blob_path(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
state = _make_typed_state()
|
||||
envelope = state.to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=0,
|
||||
)
|
||||
assert "blob_id" in envelope.payload["video"]
|
||||
assert "blob_id" in envelope.payload["audio"]
|
||||
assert "frames_b64" not in envelope.payload["video"]
|
||||
assert "safetensors_b64" not in envelope.payload["audio"]
|
||||
assert len(blob_store) == 2
|
||||
|
||||
def test_blob_roundtrip_reconstructs_tensors(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
original = _make_typed_state()
|
||||
envelope = original.to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=0,
|
||||
)
|
||||
restored = LTX2ContinuationState.from_continuation_state(
|
||||
envelope, blob_store=blob_store)
|
||||
assert restored.video_frames is not None
|
||||
assert len(restored.video_frames) == len(original.video_frames)
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, original.audio_latents)
|
||||
|
||||
def test_blob_id_held_when_store_unavailable(self):
|
||||
"""Deserializing without a blob store preserves the blob id so
|
||||
the caller can fetch it later."""
|
||||
blob_store = InMemoryBlobStore()
|
||||
envelope = _make_typed_state().to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=0,
|
||||
)
|
||||
blob_id_video = envelope.payload["video"]["blob_id"]
|
||||
blob_id_audio = envelope.payload["audio"]["blob_id"]
|
||||
|
||||
restored = LTX2ContinuationState.from_continuation_state(envelope)
|
||||
assert restored.video_frames is None
|
||||
assert restored.video_frames_blob_id == blob_id_video
|
||||
assert restored.audio_latents is None
|
||||
assert restored.audio_latents_blob_id == blob_id_audio
|
||||
|
||||
def test_large_threshold_keeps_payload_inline(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
envelope = _make_typed_state().to_continuation_state(
|
||||
blob_store=blob_store,
|
||||
inline_threshold_bytes=10 * 1024 * 1024, # 10 MiB
|
||||
)
|
||||
assert "frames_b64" in envelope.payload["video"]
|
||||
assert "safetensors_b64" in envelope.payload["audio"]
|
||||
assert len(blob_store) == 0
|
||||
|
||||
|
||||
class TestValidation:
|
||||
"""Invalid payloads error cleanly."""
|
||||
|
||||
def test_wrong_kind_rejected(self):
|
||||
envelope = ContinuationState(kind="longcat.v1", payload={})
|
||||
with pytest.raises(ValueError, match="Expected ContinuationState.kind"):
|
||||
LTX2ContinuationState.from_continuation_state(envelope)
|
||||
|
||||
def test_unsupported_schema_version_rejected(self):
|
||||
envelope = ContinuationState(
|
||||
kind=LTX2_CONTINUATION_KIND,
|
||||
payload={"schema_version": 999},
|
||||
)
|
||||
with pytest.raises(ValueError,
|
||||
match="Unsupported LTX-2 continuation schema"):
|
||||
LTX2ContinuationState.from_continuation_state(envelope)
|
||||
|
||||
def test_non_png_frame_rejected(self):
|
||||
state = LTX2ContinuationState(
|
||||
video_frames=[np.ones((64, 64, 3), dtype=np.float32)],
|
||||
)
|
||||
with pytest.raises(ValueError, match="uint8 HxWx3"):
|
||||
state.to_continuation_state()
|
||||
|
||||
|
||||
class TestCompatLayerWireUp:
|
||||
"""The public compat layer accepts request.state without reverting
|
||||
to NotImplementedError and attaches it to the SamplingParam path."""
|
||||
|
||||
def test_request_with_state_passes_through(self, tmp_path):
|
||||
# PR 7 removes the NotImplementedError for request.state; build a
|
||||
# minimal GenerationRequest carrying an LTX-2 state and make sure
|
||||
# the public boundary accepts it.
|
||||
from fastvideo.api.compat import (
|
||||
normalize_generation_request,
|
||||
_validate_continuation_state,
|
||||
)
|
||||
envelope = _make_typed_state().to_continuation_state()
|
||||
request = GenerationRequest(
|
||||
prompt="test",
|
||||
state=envelope,
|
||||
)
|
||||
normalized = normalize_generation_request(request)
|
||||
_validate_continuation_state(normalized.state)
|
||||
|
||||
def test_unknown_kind_rejected_at_boundary(self):
|
||||
from fastvideo.api.compat import _validate_continuation_state
|
||||
with pytest.raises(ValueError, match="Unknown ContinuationState kind"):
|
||||
_validate_continuation_state(
|
||||
ContinuationState(kind="mystery.v1", payload={}))
|
||||
|
||||
def test_empty_kind_rejected_at_boundary(self):
|
||||
from fastvideo.api.compat import _validate_continuation_state
|
||||
with pytest.raises(ValueError, match="non-empty string"):
|
||||
_validate_continuation_state(
|
||||
ContinuationState(kind="", payload={}))
|
||||
|
||||
def test_output_return_state_flag(self):
|
||||
request = GenerationRequest(
|
||||
prompt="x",
|
||||
output=OutputConfig(return_state=True),
|
||||
)
|
||||
# The typed public surface exposes the flag directly.
|
||||
assert request.output.return_state is True
|
||||
@@ -4,14 +4,8 @@
|
||||
Mirrors the ``load_kwargs`` dict that the FastVideo-internal
|
||||
``ui/ltx2-streaming/server/gpu_pool.py`` passes to
|
||||
``VideoGenerator.from_pretrained(**load_kwargs)`` and asserts that the
|
||||
public typed ``GeneratorConfig`` surface (introduced across PRs 0-6)
|
||||
can represent it end-to-end, with no fields silently falling through
|
||||
to ``pipeline.experimental``.
|
||||
|
||||
This is the parity guard PR 7.6 depends on: the public gpu_pool
|
||||
upstream must be able to construct a typed ``GeneratorConfig`` without
|
||||
knowing any legacy LTX-2 kwarg name, and downstream Dynamo
|
||||
(``FastVideoArgGroup``) must be able to do the same.
|
||||
public typed ``GeneratorConfig`` surface can represent it end-to-end,
|
||||
with no fields silently falling through to ``pipeline.experimental``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -27,19 +21,10 @@ from fastvideo.api.compat import (
|
||||
|
||||
# Mirrors FastVideo-internal/ui/ltx2-streaming/server/gpu_pool.py
|
||||
# :lines 233-260 (load_kwargs constructed for VideoGenerator.from_pretrained).
|
||||
#
|
||||
# One item from gpu_pool.py's load_kwargs is deliberately excluded:
|
||||
# - ``pipeline_config=<PipelineConfig instance>`` — an opaque Python
|
||||
# object; internal mutates it in place (``dit_config.quant_config =
|
||||
# FP4Config()``). The typed path for quantization is tracked in
|
||||
# "Known Technical Debt" in PR plan.md; ``pipeline_config`` as an
|
||||
# instance legitimately belongs in ``pipeline.experimental``.
|
||||
#
|
||||
# ``enable_torch_compile_text_encoder`` IS included below: its typed
|
||||
# home is ``CompileConfig.text_encoder_enabled`` (added post-review).
|
||||
# The legacy ``FastVideoArgs`` path does not yet consume it; the
|
||||
# realtime runtime (PR 7.6) reads it off the kwargs dict before
|
||||
# FastVideoArgs filtering.
|
||||
# Skipped here because they are opaque Python objects that legitimately
|
||||
# belong in experimental:
|
||||
# - pipeline_config=<PipelineConfig instance>
|
||||
# - enable_torch_compile_text_encoder (not in public FastVideoArgs)
|
||||
GPU_POOL_LOAD_KWARGS = {
|
||||
"config_model_path": "/models/ltx2-distilled/config",
|
||||
"num_gpus": 1,
|
||||
@@ -57,7 +42,6 @@ GPU_POOL_LOAD_KWARGS = {
|
||||
"ltx2_refine_guidance_scale": 1.0,
|
||||
"ltx2_refine_add_noise": True,
|
||||
"enable_torch_compile": True,
|
||||
"enable_torch_compile_text_encoder": True,
|
||||
"torch_compile_kwargs": {
|
||||
"backend": "inductor",
|
||||
"fullgraph": True,
|
||||
@@ -94,7 +78,6 @@ class TestGpuPoolForwardTranslation:
|
||||
def test_compile_config_typed_fields_extracted(self, config) -> None:
|
||||
compile_config = config.engine.compile
|
||||
assert compile_config.enabled is True
|
||||
assert compile_config.text_encoder_enabled is True
|
||||
assert compile_config.backend == "inductor"
|
||||
assert compile_config.fullgraph is True
|
||||
assert compile_config.mode == "max-autotune-no-cudagraphs"
|
||||
@@ -133,11 +116,9 @@ class TestGpuPoolForwardTranslation:
|
||||
|
||||
class TestGpuPoolReverseTranslation:
|
||||
"""typed GeneratorConfig -> FastVideoArgs kwargs reproduces the
|
||||
original gpu_pool flat-kwarg shape.
|
||||
|
||||
This is what lets PR 7.6 wire the public ``gpu_pool`` through
|
||||
``generator_config_to_fastvideo_args`` without the runtime noticing.
|
||||
"""
|
||||
original gpu_pool flat-kwarg shape, so callers can wire a public
|
||||
``gpu_pool`` through ``generator_config_to_fastvideo_args`` without
|
||||
the runtime noticing."""
|
||||
|
||||
@pytest.fixture
|
||||
def args_kwargs(self, monkeypatch):
|
||||
@@ -187,12 +168,6 @@ class TestGpuPoolReverseTranslation:
|
||||
def test_vae_tiling_reemitted_with_legacy_name(self, args_kwargs) -> None:
|
||||
assert args_kwargs["ltx2_vae_tiling"] is False
|
||||
|
||||
def test_text_encoder_compile_reemitted(self, args_kwargs) -> None:
|
||||
# Present in the captured kwargs dict even though
|
||||
# ``FastVideoArgs.from_kwargs`` will filter it out — realtime
|
||||
# runtime upstream (PR 7.6) reads it off this dict.
|
||||
assert args_kwargs["enable_torch_compile_text_encoder"] is True
|
||||
|
||||
def test_no_stray_refine_dict(self, args_kwargs) -> None:
|
||||
"""preset_overrides.refine must flatten to ltx2_refine_* kwargs
|
||||
rather than landing as a nested ``refine`` kwarg that
|
||||
@@ -213,9 +188,7 @@ class TestRefineFlattenCoversAllTypedFields:
|
||||
)
|
||||
from fastvideo.api.schema import GeneratorConfig, PipelineSelection
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
REFINE_FLAT_KEYS, )
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
@@ -240,9 +213,7 @@ class TestRefineFlattenCoversAllTypedFields:
|
||||
"image_crf": 18,
|
||||
"video_position_offset_sec": 2.5,
|
||||
}
|
||||
all_fields = (refine_preset_override_fields()
|
||||
| refine_stage_override_fields())
|
||||
assert set(refine_payload) == all_fields, (
|
||||
assert set(refine_payload) == REFINE_FLAT_KEYS, (
|
||||
"payload must cover every typed field to exercise the flatten loop")
|
||||
|
||||
config = GeneratorConfig(
|
||||
|
||||
@@ -9,9 +9,9 @@ from fastvideo.api.presets import get_preset, validate_stage_overrides
|
||||
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
|
||||
LTX2RefinePresetOverride,
|
||||
LTX2RefineStageOverride,
|
||||
REFINE_PRESET_OVERRIDE_FIELDS,
|
||||
REFINE_STAGE_OVERRIDE_FIELDS,
|
||||
refine_override_to_dict,
|
||||
refine_preset_override_fields,
|
||||
refine_stage_override_fields,
|
||||
)
|
||||
|
||||
|
||||
@@ -57,7 +57,7 @@ class TestRefineStageOverrideDataclass:
|
||||
}
|
||||
|
||||
def test_fields_accessor_matches_dataclass(self) -> None:
|
||||
assert refine_stage_override_fields() == frozenset({
|
||||
assert REFINE_STAGE_OVERRIDE_FIELDS == frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
"image_crf",
|
||||
@@ -86,7 +86,7 @@ class TestRefinePresetOverrideDataclass:
|
||||
}
|
||||
|
||||
def test_fields_accessor_matches_dataclass(self) -> None:
|
||||
assert refine_preset_override_fields() == frozenset({
|
||||
assert REFINE_PRESET_OVERRIDE_FIELDS == frozenset({
|
||||
"enabled",
|
||||
"add_noise",
|
||||
})
|
||||
@@ -101,7 +101,7 @@ class TestStageOverridesMirrorPresetSchema:
|
||||
preset = get_preset("ltx2_two_stage", "ltx2")
|
||||
refine_schema = next(
|
||||
s for s in preset.stage_schemas if s.name == "refine")
|
||||
assert refine_schema.allowed_overrides == refine_stage_override_fields()
|
||||
assert refine_schema.allowed_overrides == REFINE_STAGE_OVERRIDE_FIELDS
|
||||
|
||||
def test_roundtrip_through_validate_stage_overrides(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
|
||||
@@ -113,7 +113,6 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
|
||||
},
|
||||
"compile": {
|
||||
"enabled": False,
|
||||
"text_encoder_enabled": None,
|
||||
"backend": None,
|
||||
"fullgraph": None,
|
||||
"mode": None,
|
||||
|
||||
@@ -585,40 +585,3 @@ class TestPresetCountIntegrity:
|
||||
import fastvideo.registry # noqa: F401
|
||||
names = get_all_preset_names()
|
||||
assert len(names) >= 37
|
||||
|
||||
|
||||
class TestPresetDefaultTypes:
|
||||
"""Preset ``defaults`` values must match the types on
|
||||
:class:`SamplingParam`. Assigning ``None`` to a typed-``str`` field
|
||||
(e.g. ``negative_prompt``) breaks downstream stages that assert the
|
||||
runtime type — see the CFG branch in
|
||||
``pipelines/stages/text_encoding.py:81``."""
|
||||
|
||||
def test_ltx2_cfg_defaults_are_off(self) -> None:
|
||||
"""SamplingParam's LTX-2 CFG class defaults must be 1.0 (CFG
|
||||
off). ``ForwardBatch.__post_init__`` force-enables
|
||||
``do_classifier_free_guidance`` when either
|
||||
``ltx2_cfg_scale_video`` or ``ltx2_cfg_scale_audio`` is != 1.0,
|
||||
so any non-1.0 default silently forces CFG on for every model
|
||||
family that doesn't explicitly override these fields. Guard
|
||||
against the regression that surfaced as the TurboDiffusion I2V
|
||||
SSIM crash (``text_encoding.py:81`` assertion on
|
||||
``negative_prompt``)."""
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
sp = SamplingParam()
|
||||
assert sp.ltx2_cfg_scale_video == 1.0
|
||||
assert sp.ltx2_cfg_scale_audio == 1.0
|
||||
|
||||
def test_no_preset_sets_negative_prompt_to_none(self) -> None:
|
||||
import fastvideo.registry # noqa: F401
|
||||
from fastvideo.api.presets import _PRESET_REGISTRY
|
||||
offenders = [
|
||||
f"{preset.model_family}/{preset.name}"
|
||||
for preset in _PRESET_REGISTRY.values()
|
||||
if preset.defaults.get("negative_prompt", "") is None
|
||||
]
|
||||
assert not offenders, (
|
||||
"These presets set negative_prompt=None, which violates "
|
||||
"SamplingParam.negative_prompt's typed str contract and "
|
||||
"crashes the CFG path in text_encoding. Use \"\" instead:\n"
|
||||
+ "\n".join(f" - {p}" for p in offenders))
|
||||
|
||||
@@ -71,7 +71,13 @@ def _get_extra_dataclass_fields(
|
||||
continue
|
||||
for _, modname, is_pkg in pkgutil.walk_packages(
|
||||
package.__path__, prefix=f"{package_name}."):
|
||||
if modname.endswith(".__pycache__"):
|
||||
# Flat ``configs.pipelines.<family>`` modules carry the config
|
||||
# directly; colocated ``basic.<family>.pipeline_configs``
|
||||
# submodules do too. Everything else under ``basic`` is heavy
|
||||
# model code we don't need to import for a schema check.
|
||||
basename = modname.rsplit(".", 1)[-1]
|
||||
is_flat = modname.startswith("fastvideo.configs.pipelines.")
|
||||
if not is_flat and basename != "pipeline_configs":
|
||||
continue
|
||||
module = importlib.import_module(modname)
|
||||
for obj in vars(module).values():
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -1,154 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Protocol schema tests for the streaming server.
|
||||
|
||||
Covers:
|
||||
|
||||
* accepted client messages parse into the correct discriminated model
|
||||
* unknown ``type`` values raise validation errors
|
||||
* server-side messages serialize to the expected wire shape
|
||||
* continuation_state field on session_init_v2 carries through
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
ContinuationStateSnapshot,
|
||||
ErrorMessage,
|
||||
GpuAssigned,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
|
||||
|
||||
class TestClientMessageParsing:
|
||||
|
||||
def test_session_init_v2_minimal(self):
|
||||
parsed = parse_client_message({"type": "session_init_v2"})
|
||||
assert isinstance(parsed, SessionInitV2)
|
||||
assert parsed.curated_prompts == []
|
||||
assert parsed.stream_mode == "av_fmp4"
|
||||
|
||||
def test_session_init_v2_full(self):
|
||||
raw = {
|
||||
"type": "session_init_v2",
|
||||
"client_id": "client-1",
|
||||
"preset": "ltx2_two_stage",
|
||||
"preset_label": "2x refine",
|
||||
"curated_prompts": ["a fox", "a deer"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
"single_clip_mode": True,
|
||||
"stream_mode": "av_fmp4",
|
||||
"continuation_state": {
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, "segment_index": 2},
|
||||
},
|
||||
}
|
||||
parsed = parse_client_message(raw)
|
||||
assert isinstance(parsed, SessionInitV2)
|
||||
assert parsed.preset == "ltx2_two_stage"
|
||||
assert parsed.curated_prompts == ["a fox", "a deer"]
|
||||
assert parsed.continuation_state["kind"] == "ltx2.v1"
|
||||
|
||||
def test_segment_prompt_source(self):
|
||||
parsed = parse_client_message({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "hello world",
|
||||
"source": "curated",
|
||||
"seed": 7,
|
||||
})
|
||||
assert isinstance(parsed, SegmentPromptSource)
|
||||
assert parsed.source == "curated"
|
||||
assert parsed.seed == 7
|
||||
|
||||
def test_snapshot_state(self):
|
||||
parsed = parse_client_message({"type": "snapshot_state"})
|
||||
assert isinstance(parsed, SnapshotState)
|
||||
|
||||
def test_unknown_type_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"type": "not_a_real_message"})
|
||||
|
||||
def test_missing_type_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"prompt": "x"})
|
||||
|
||||
def test_segment_prompt_source_requires_prompt(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"type": "segment_prompt_source"})
|
||||
|
||||
|
||||
class TestServerMessageSerialization:
|
||||
|
||||
def test_queue_status(self):
|
||||
msg = QueueStatus(position=3, queue_depth=5)
|
||||
assert msg.model_dump() == {
|
||||
"type": "queue_status",
|
||||
"position": 3,
|
||||
"queue_depth": 5,
|
||||
}
|
||||
|
||||
def test_gpu_assigned(self):
|
||||
msg = GpuAssigned(gpu_id=1, session_timeout=300)
|
||||
assert msg.model_dump()["type"] == "gpu_assigned"
|
||||
|
||||
def test_ltx2_stream_start(self):
|
||||
msg = Ltx2StreamStart(
|
||||
preset="ltx2_two_stage",
|
||||
width=1024, height=1536, fps=24, num_frames=121,
|
||||
)
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["type"] == "ltx2_stream_start"
|
||||
assert dumped["width"] == 1024
|
||||
|
||||
def test_ltx2_segment_start(self):
|
||||
msg = Ltx2SegmentStart(
|
||||
segment_idx=0,
|
||||
prompt="a fox",
|
||||
total_steps=8,
|
||||
)
|
||||
assert msg.model_dump()["segment_idx"] == 0
|
||||
|
||||
def test_step_complete(self):
|
||||
msg = StepComplete(segment_idx=0, step=1, total_steps=8)
|
||||
assert msg.model_dump()["stage"] == "denoise"
|
||||
|
||||
def test_media_init_has_mode(self):
|
||||
msg = MediaInit(segment_idx=0, stream_id="abc")
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["mode"] == "av_fmp4"
|
||||
assert "avc1" in dumped["mime"]
|
||||
|
||||
def test_media_segment_complete(self):
|
||||
msg = MediaSegmentComplete(
|
||||
segment_idx=0, stream_id="abc", chunks=4,
|
||||
)
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["chunks"] == 4
|
||||
|
||||
def test_ltx2_segment_complete(self):
|
||||
msg = Ltx2SegmentComplete(segment_idx=0, generation_time_ms=1234.5)
|
||||
assert msg.model_dump()["generation_time_ms"] == 1234.5
|
||||
|
||||
def test_error_message_code_restricted(self):
|
||||
with pytest.raises(ValidationError):
|
||||
ErrorMessage(code="not_a_code", message="x")
|
||||
|
||||
def test_continuation_state_snapshot(self):
|
||||
msg = ContinuationStateSnapshot(state={
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1},
|
||||
})
|
||||
assert msg.model_dump()["state"]["kind"] == "ltx2.v1"
|
||||
@@ -1,237 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""End-to-end WebSocket smoke for the streaming server skeleton.
|
||||
|
||||
Uses a mock generator so these tests run CPU-only (no GPU, no model
|
||||
weights). Skips the fMP4 assertions when ``ffmpeg`` is missing.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("starlette")
|
||||
from starlette.testclient import TestClient # noqa: E402
|
||||
|
||||
from fastvideo.api.schema import ( # noqa: E402
|
||||
ContinuationState,
|
||||
GeneratorConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
StreamingConfig,
|
||||
GenerationRequest,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.server import build_app # noqa: E402
|
||||
|
||||
|
||||
_FFMPEG_AVAILABLE = shutil.which("ffmpeg") is not None
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MockGenerator:
|
||||
width: int = 64
|
||||
height: int = 64
|
||||
fps: int = 12
|
||||
num_frames: int = 12
|
||||
return_state: bool = True
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
frames = [
|
||||
np.full((self.height, self.width, 3), i * 5, dtype=np.uint8)
|
||||
for i in range(self.num_frames)
|
||||
]
|
||||
state = (ContinuationState(
|
||||
kind="ltx2.v1",
|
||||
payload={
|
||||
"schema_version": 1,
|
||||
"segment_index": 0,
|
||||
"source_prompt": request.prompt,
|
||||
},
|
||||
) if self.return_state else None)
|
||||
return {
|
||||
"frames": frames,
|
||||
"audio_sample_rate": 24000,
|
||||
"state": state,
|
||||
}
|
||||
|
||||
|
||||
def _build_serve_config() -> ServeConfig:
|
||||
return ServeConfig(
|
||||
generator=GeneratorConfig(model_path="/models/fake"),
|
||||
default_request=GenerationRequest(
|
||||
sampling=SamplingConfig(
|
||||
num_frames=12,
|
||||
height=64,
|
||||
width=64,
|
||||
fps=12,
|
||||
num_inference_steps=1,
|
||||
),
|
||||
),
|
||||
streaming=StreamingConfig(
|
||||
session_timeout_seconds=60,
|
||||
generation_segment_cap=2,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _build_client() -> tuple[TestClient, _MockGenerator]:
|
||||
generator = _MockGenerator()
|
||||
app = build_app(_build_serve_config(), generator)
|
||||
return TestClient(app), generator
|
||||
|
||||
|
||||
class TestHealth:
|
||||
|
||||
def test_health_endpoint_reports_stream_mode(self):
|
||||
client, _ = _build_client()
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["stream_mode"] == "av_fmp4"
|
||||
assert body["sessions"] == 0
|
||||
|
||||
|
||||
class TestSessionHandshake:
|
||||
|
||||
def test_rejects_non_session_init_opening_frame(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "segment_prompt_source", "prompt": "x"})
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
assert err["code"] == "invalid_message"
|
||||
|
||||
def test_rejects_unknown_message_on_init(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "not_a_message"})
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
|
||||
def test_emits_queue_and_gpu_assigned_on_valid_init(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({
|
||||
"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox"],
|
||||
})
|
||||
assert ws.receive_json()["type"] == "queue_status"
|
||||
assert ws.receive_json()["type"] == "gpu_assigned"
|
||||
assert ws.receive_json()["type"] == "ltx2_stream_start"
|
||||
|
||||
def test_init_hydrates_continuation_state(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({
|
||||
"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"continuation_state": {
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, "segment_index": 3},
|
||||
},
|
||||
})
|
||||
# Drain handshake frames
|
||||
ws.receive_json() # queue_status
|
||||
ws.receive_json() # gpu_assigned
|
||||
ws.receive_json() # ltx2_stream_start
|
||||
# Ask the server for the state back; it should echo what we sent.
|
||||
ws.send_json({"type": "snapshot_state"})
|
||||
snap = ws.receive_json()
|
||||
assert snap["type"] == "continuation_state_snapshot"
|
||||
assert snap["state"]["kind"] == "ltx2.v1"
|
||||
assert snap["state"]["payload"]["segment_index"] == 3
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _FFMPEG_AVAILABLE, reason="ffmpeg not installed")
|
||||
class TestSegmentFlow:
|
||||
|
||||
def test_segment_generates_media_init_plus_complete(self):
|
||||
client, generator = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage"})
|
||||
for _ in range(3):
|
||||
ws.receive_json() # queue_status + gpu_assigned + stream_start
|
||||
|
||||
ws.send_json({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "a test segment",
|
||||
"num_inference_steps": 1,
|
||||
})
|
||||
start = ws.receive_json()
|
||||
assert start["type"] == "ltx2_segment_start"
|
||||
assert start["segment_idx"] == 0
|
||||
step = ws.receive_json()
|
||||
assert step["type"] == "step_complete"
|
||||
media_init = ws.receive_json()
|
||||
assert media_init["type"] == "media_init"
|
||||
# Then one or more binary frames until media_segment_complete.
|
||||
saw_binary = False
|
||||
while True:
|
||||
msg = ws.receive()
|
||||
if "bytes" in msg and msg["bytes"]:
|
||||
saw_binary = True
|
||||
continue
|
||||
parsed = _as_json(msg)
|
||||
if parsed is None:
|
||||
continue
|
||||
if parsed["type"] == "media_segment_complete":
|
||||
break
|
||||
assert saw_binary
|
||||
final = ws.receive_json()
|
||||
assert final["type"] == "ltx2_segment_complete"
|
||||
assert final["segment_idx"] == 0
|
||||
|
||||
|
||||
class TestContinuationStatePersistence:
|
||||
|
||||
def test_snapshot_after_segment_carries_generator_state(self):
|
||||
if not _FFMPEG_AVAILABLE:
|
||||
pytest.skip("ffmpeg not installed")
|
||||
client, generator = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage"})
|
||||
for _ in range(3):
|
||||
ws.receive_json()
|
||||
ws.send_json({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "a cat",
|
||||
"num_inference_steps": 1,
|
||||
})
|
||||
_drain_until(ws, "ltx2_segment_complete")
|
||||
ws.send_json({"type": "snapshot_state"})
|
||||
snap = ws.receive_json()
|
||||
assert snap["type"] == "continuation_state_snapshot"
|
||||
assert snap["state"]["kind"] == "ltx2.v1"
|
||||
assert snap["state"]["payload"]["source_prompt"] == "a cat"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
def _drain_until(ws, target_type: str) -> dict[str, Any]:
|
||||
while True:
|
||||
msg = ws.receive()
|
||||
if "text" in msg and msg["text"]:
|
||||
import json
|
||||
|
||||
parsed = json.loads(msg["text"])
|
||||
if parsed.get("type") == target_type:
|
||||
return parsed
|
||||
# skip binary / other
|
||||
|
||||
|
||||
def _as_json(msg: dict[str, Any]) -> dict[str, Any] | None:
|
||||
if "text" not in msg or not msg["text"]:
|
||||
return None
|
||||
import json
|
||||
|
||||
return json.loads(msg["text"])
|
||||
@@ -1,133 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Session lifecycle tests."""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
|
||||
|
||||
class TestSessionStateMachine:
|
||||
|
||||
def test_starts_initializing(self):
|
||||
s = Session()
|
||||
assert s.state is SessionState.INITIALIZING
|
||||
|
||||
def test_legal_sequence(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.COMPLETE)
|
||||
assert s.state is SessionState.COMPLETE
|
||||
|
||||
def test_active_self_loop_allowed(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.ACTIVE) # re-asserting is fine
|
||||
assert s.state is SessionState.ACTIVE
|
||||
|
||||
def test_illegal_backwards_transition(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.INITIALIZING)
|
||||
|
||||
def test_cannot_leave_terminal_state(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.COMPLETE)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.ACTIVE)
|
||||
|
||||
def test_error_terminal(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.ERROR)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.ACTIVE)
|
||||
|
||||
def test_transition_updates_activity(self):
|
||||
s = Session()
|
||||
prior = s.last_activity
|
||||
time.sleep(0.001)
|
||||
s.transition(SessionState.QUEUED)
|
||||
assert s.last_activity > prior
|
||||
|
||||
def test_segment_cap(self):
|
||||
s = Session()
|
||||
s.segment_idx = 5
|
||||
assert s.segment_cap_reached(5) is True
|
||||
assert s.segment_cap_reached(6) is False
|
||||
|
||||
|
||||
class TestSessionManager:
|
||||
|
||||
def test_create_assigns_unique_ids(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=2)
|
||||
a = mgr.create()
|
||||
b = mgr.create()
|
||||
assert a.id != b.id
|
||||
assert len(mgr) == 2
|
||||
|
||||
def test_max_sessions_enforced(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
|
||||
mgr.create()
|
||||
with pytest.raises(SessionRejected):
|
||||
mgr.create()
|
||||
|
||||
def test_close_releases_slot(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
|
||||
s = mgr.create()
|
||||
mgr.close(s.id)
|
||||
assert len(mgr) == 0
|
||||
# Now can create again.
|
||||
mgr.create()
|
||||
|
||||
def test_reap_timed_out_flags_idle_sessions(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
|
||||
s = mgr.create()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.last_activity = time.monotonic() - 10 # 10s ago, past the 1s budget
|
||||
dead = mgr.reap_timed_out()
|
||||
assert s.id in dead
|
||||
|
||||
def test_reap_skips_terminal_states(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
|
||||
s = mgr.create()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.ERROR)
|
||||
s.last_activity = time.monotonic() - 10
|
||||
assert s.id not in mgr.reap_timed_out()
|
||||
|
||||
def test_active_sessions_filter(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=4)
|
||||
a = mgr.create()
|
||||
a.transition(SessionState.QUEUED)
|
||||
a.transition(SessionState.GPU_BINDING)
|
||||
a.transition(SessionState.ACTIVE)
|
||||
b = mgr.create() # INITIALIZING
|
||||
assert mgr.active_sessions() == [a]
|
||||
assert b not in mgr.active_sessions()
|
||||
@@ -1,90 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the session init-image persistence helper."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image,
|
||||
)
|
||||
|
||||
|
||||
def _png_bytes(size: tuple[int, int] = (64, 64)) -> bytes:
|
||||
buffer = io.BytesIO()
|
||||
Image.new("RGB", size, color=(10, 20, 30)).save(buffer, format="PNG")
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
class TestPersistSessionInitImage:
|
||||
|
||||
def test_none_payload_returns_none(self):
|
||||
assert persist_session_init_image(None) is None
|
||||
assert persist_session_init_image({}) is None
|
||||
|
||||
def test_non_object_payload_rejected(self):
|
||||
with pytest.raises(ValueError):
|
||||
persist_session_init_image("not-a-dict")
|
||||
|
||||
def test_png_payload_persists(self, tmp_path):
|
||||
data = _png_bytes()
|
||||
image = persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"name": "ref.png",
|
||||
"data": base64.b64encode(data).decode("ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
assert image is not None
|
||||
assert os.path.exists(image.path)
|
||||
assert image.mime == "image/png"
|
||||
assert image.path.endswith(".png")
|
||||
with open(image.path, "rb") as f:
|
||||
assert f.read() == data
|
||||
|
||||
def test_unknown_mime_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="mime"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/bmp",
|
||||
"data": "ignored",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_bad_base64_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="base64"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": "not!base64!",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_empty_data_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="empty"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": "",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_display_name_sanitized(self, tmp_path):
|
||||
image = persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"name": "../evil/../name.png",
|
||||
"data": base64.b64encode(_png_bytes()).decode("ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
assert image is not None
|
||||
assert image.display_name == "name.png"
|
||||
|
||||
def test_oversize_rejected(self, tmp_path):
|
||||
from fastvideo.entrypoints.streaming import session_init_image as mod
|
||||
|
||||
original = mod._MAX_IMAGE_BYTES
|
||||
mod._MAX_IMAGE_BYTES = 100
|
||||
try:
|
||||
with pytest.raises(ValueError, match="limit"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": base64.b64encode(_png_bytes((512, 512))).decode(
|
||||
"ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
finally:
|
||||
mod._MAX_IMAGE_BYTES = original
|
||||
@@ -1,186 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the streaming SessionStore and BlobStore.
|
||||
|
||||
Covers:
|
||||
|
||||
* ``store`` / ``snapshot`` / ``drop`` lifecycle for the in-memory store
|
||||
* ``hydrate`` with and without an explicit session id
|
||||
* blob store insert / get / drop semantics
|
||||
* thread-safety under concurrent writes (smoke)
|
||||
* round-trip a LTX-2 continuation through snapshot + hydrate across a
|
||||
session boundary (the "export and resume" flow the PR plan calls out)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
BlobStore,
|
||||
InMemoryBlobStore,
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.pipelines.basic.ltx2.continuation import (
|
||||
LTX2_CONTINUATION_KIND,
|
||||
LTX2ContinuationState,
|
||||
)
|
||||
|
||||
|
||||
class TestInMemoryBlobStore:
|
||||
|
||||
def test_is_blob_store(self):
|
||||
assert isinstance(InMemoryBlobStore(), BlobStore)
|
||||
|
||||
def test_put_then_get_returns_same_bytes(self):
|
||||
store = InMemoryBlobStore()
|
||||
blob_id = store.put(b"hello")
|
||||
assert store.get(blob_id) == b"hello"
|
||||
|
||||
def test_put_returns_distinct_ids(self):
|
||||
store = InMemoryBlobStore()
|
||||
id_a = store.put(b"a")
|
||||
id_b = store.put(b"b")
|
||||
assert id_a != id_b
|
||||
|
||||
def test_get_missing_raises_keyerror(self):
|
||||
store = InMemoryBlobStore()
|
||||
with pytest.raises(KeyError):
|
||||
store.get("nonexistent")
|
||||
|
||||
def test_drop_removes_blob(self):
|
||||
store = InMemoryBlobStore()
|
||||
blob_id = store.put(b"payload")
|
||||
store.drop(blob_id)
|
||||
assert blob_id not in store
|
||||
with pytest.raises(KeyError):
|
||||
store.get(blob_id)
|
||||
|
||||
def test_drop_missing_is_noop(self):
|
||||
store = InMemoryBlobStore()
|
||||
store.drop("not-there") # no raise
|
||||
|
||||
def test_contains(self):
|
||||
store = InMemoryBlobStore()
|
||||
blob_id = store.put(b"x")
|
||||
assert blob_id in store
|
||||
assert "other" not in store
|
||||
|
||||
|
||||
class TestInMemorySessionStore:
|
||||
|
||||
def test_is_session_store(self):
|
||||
assert isinstance(InMemorySessionStore(), SessionStore)
|
||||
|
||||
def test_store_then_snapshot(self):
|
||||
store = InMemorySessionStore()
|
||||
state = ContinuationState(kind="ltx2.v1", payload={"x": 1})
|
||||
store.store("sess-1", state)
|
||||
assert store.snapshot("sess-1") is state
|
||||
|
||||
def test_snapshot_missing_returns_none(self):
|
||||
store = InMemorySessionStore()
|
||||
assert store.snapshot("missing") is None
|
||||
|
||||
def test_store_overwrites_prior_state(self):
|
||||
store = InMemorySessionStore()
|
||||
first = ContinuationState(kind="ltx2.v1", payload={"v": 1})
|
||||
second = ContinuationState(kind="ltx2.v1", payload={"v": 2})
|
||||
store.store("s", first)
|
||||
store.store("s", second)
|
||||
assert store.snapshot("s").payload["v"] == 2
|
||||
|
||||
def test_hydrate_assigns_new_session_id(self):
|
||||
store = InMemorySessionStore()
|
||||
state = ContinuationState(kind="ltx2.v1", payload={})
|
||||
sid = store.hydrate(state)
|
||||
assert sid
|
||||
assert store.snapshot(sid) is state
|
||||
|
||||
def test_hydrate_with_explicit_session_id(self):
|
||||
store = InMemorySessionStore()
|
||||
state = ContinuationState(kind="ltx2.v1", payload={})
|
||||
sid = store.hydrate(state, session_id="pinned-id")
|
||||
assert sid == "pinned-id"
|
||||
assert store.snapshot("pinned-id") is state
|
||||
|
||||
def test_drop_forgets_session(self):
|
||||
store = InMemorySessionStore()
|
||||
store.store("s", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
store.drop("s")
|
||||
assert store.snapshot("s") is None
|
||||
assert "s" not in store
|
||||
|
||||
def test_iter_yields_session_ids(self):
|
||||
store = InMemorySessionStore()
|
||||
store.store("a", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
store.store("b", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
assert sorted(store) == ["a", "b"]
|
||||
|
||||
def test_len(self):
|
||||
store = InMemorySessionStore()
|
||||
assert len(store) == 0
|
||||
store.store("x", ContinuationState(kind="ltx2.v1", payload={}))
|
||||
assert len(store) == 1
|
||||
|
||||
def test_concurrent_store_is_safe(self):
|
||||
"""Smoke-check the lock: 200 parallel stores settle to 200 ids."""
|
||||
store = InMemorySessionStore()
|
||||
|
||||
def write(i: int) -> None:
|
||||
store.store(
|
||||
f"s-{i}",
|
||||
ContinuationState(kind="ltx2.v1", payload={"i": i}),
|
||||
)
|
||||
|
||||
threads = [threading.Thread(target=write, args=(i,)) for i in range(200)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
assert len(store) == 200
|
||||
|
||||
|
||||
class TestSnapshotHydrateRoundTrip:
|
||||
"""Session boundary: snapshot + hydrate preserves the full LTX-2 state."""
|
||||
|
||||
def test_end_to_end_ltx2_session_migration(self):
|
||||
blob_store = InMemoryBlobStore()
|
||||
sessions = InMemorySessionStore()
|
||||
|
||||
typed = LTX2ContinuationState(
|
||||
segment_index=4,
|
||||
video_frames=[
|
||||
np.full((32, 32, 3), i * 5, dtype=np.uint8) for i in range(3)
|
||||
],
|
||||
audio_latents=torch.randn(1, 4, 8, 32, dtype=torch.float32),
|
||||
audio_sample_rate=24000,
|
||||
audio_conditioning_num_frames=5,
|
||||
video_position_offset_sec=0.25,
|
||||
)
|
||||
envelope = typed.to_continuation_state(blob_store=blob_store)
|
||||
|
||||
sessions.store("session-a", envelope)
|
||||
snapshot = sessions.snapshot("session-a")
|
||||
assert snapshot is not None
|
||||
assert snapshot.kind == LTX2_CONTINUATION_KIND
|
||||
|
||||
# Simulate a migration: drop the first session, hydrate a new one
|
||||
# from the snapshot, and reconstruct the typed state.
|
||||
sessions.drop("session-a")
|
||||
new_sid = sessions.hydrate(snapshot)
|
||||
assert new_sid != "session-a"
|
||||
rebuilt = sessions.snapshot(new_sid)
|
||||
assert rebuilt is snapshot
|
||||
|
||||
restored = LTX2ContinuationState.from_continuation_state(
|
||||
rebuilt, blob_store=blob_store)
|
||||
assert restored.segment_index == typed.segment_index
|
||||
assert restored.audio_sample_rate == typed.audio_sample_rate
|
||||
torch.testing.assert_close(
|
||||
restored.audio_latents, typed.audio_latents)
|
||||
assert len(restored.video_frames) == len(typed.video_frames)
|
||||
@@ -1,99 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the fMP4 encoder.
|
||||
|
||||
These tests require ``ffmpeg`` on PATH. Skip when missing so the suite
|
||||
stays CPU/CI friendly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import shutil
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
shutil.which("ffmpeg") is None,
|
||||
reason="ffmpeg not installed",
|
||||
)
|
||||
|
||||
|
||||
def _frame(width: int, height: int, value: int = 128) -> np.ndarray:
|
||||
return np.full((height, width, 3), value, dtype=np.uint8)
|
||||
|
||||
|
||||
def test_encoder_emits_init_then_media_chunks():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
chunks: list[FragmentedMP4Chunk] = []
|
||||
async with enc:
|
||||
frames = [_frame(64, 64, v) for v in range(4, 28)]
|
||||
async for chunk in enc.encode(frames):
|
||||
chunks.append(chunk)
|
||||
assert len(chunks) > 0
|
||||
assert chunks[0].kind == "init"
|
||||
assert all(c.stream_id == enc.stream_id for c in chunks)
|
||||
assert all(c.segment_idx == 0 for c in chunks)
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_init_chunk_is_fmp4():
|
||||
"""The first chunk must contain the ``ftyp`` box (fMP4 init segment)."""
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
first_chunk = None
|
||||
async with enc:
|
||||
async for chunk in enc.encode([_frame(64, 64, 20)] * 24):
|
||||
first_chunk = chunk
|
||||
break
|
||||
assert first_chunk is not None
|
||||
assert first_chunk.kind == "init"
|
||||
# Box header: 4 bytes length, 4 bytes type. "ftyp" should appear
|
||||
# near the start of the init segment.
|
||||
assert b"ftyp" in first_chunk.data[:32]
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_rejects_non_ndarray_frames():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
async with enc:
|
||||
with pytest.raises(TypeError):
|
||||
async for _ in enc.encode(["not-a-frame"]):
|
||||
pass
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_rejects_wrong_shape():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
async with enc:
|
||||
with pytest.raises(ValueError):
|
||||
async for _ in enc.encode(
|
||||
[np.zeros((64, 64, 4), dtype=np.uint8)]):
|
||||
pass
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_close_is_idempotent():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
await enc.__aenter__()
|
||||
await enc.close()
|
||||
await enc.close() # no raise
|
||||
|
||||
asyncio.run(run())
|
||||
@@ -24,10 +24,6 @@ image = (modal.Image.from_registry(
|
||||
os.environ.get("BUILDKITE_COMMIT", ""),
|
||||
"BUILDKITE_PULL_REQUEST":
|
||||
os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"BUILDKITE_BRANCH":
|
||||
os.environ.get("BUILDKITE_BRANCH", ""),
|
||||
"TEST_SCOPE":
|
||||
os.environ.get("TEST_SCOPE", ""),
|
||||
"IMAGE_VERSION":
|
||||
os.environ.get("IMAGE_VERSION", ""),
|
||||
}))
|
||||
@@ -70,21 +66,13 @@ def run_test(pytest_command: str):
|
||||
{pytest_command}
|
||||
"""
|
||||
|
||||
# result = subprocess.run(["/bin/bash", "-c", command],
|
||||
# stdout=sys.stdout,
|
||||
# stderr=sys.stderr,
|
||||
# check=False)
|
||||
|
||||
# sys.exit(result.returncode)
|
||||
|
||||
result = subprocess.run(["/bin/bash", "-c", command],
|
||||
stdout=sys.stdout,
|
||||
stderr=sys.stderr,
|
||||
check=False)
|
||||
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Test command failed with exit code {result.returncode}")
|
||||
# On success, just return — don't call sys.exit()
|
||||
sys.exit(result.returncode)
|
||||
|
||||
|
||||
@app.function(gpu="H100:1",
|
||||
image=image,
|
||||
@@ -218,7 +206,7 @@ def run_self_forcing_tests():
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test(
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs"
|
||||
)
|
||||
|
||||
|
||||
@@ -240,20 +228,12 @@ def run_lora_extraction_tests():
|
||||
timeout=1800,
|
||||
secrets=[
|
||||
modal.Secret.from_dict(
|
||||
{"HF_API_KEY": os.environ.get("HF_API_KEY", ""),
|
||||
"HF_REPO_ID": "FastVideo/performance-tracking"})
|
||||
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
|
||||
],
|
||||
volumes={
|
||||
"/root/data": model_vol,
|
||||
})
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_performance_tests():
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && "
|
||||
"export PERFORMANCE_TRACKING_ROOT='/tmp/perf-tracking' && "
|
||||
"hf auth login --token $HF_API_KEY && "
|
||||
"pytest ./fastvideo/tests/performance -vs && "
|
||||
"python ./fastvideo/tests/performance/compare_baseline.py && "
|
||||
"python ./fastvideo/tests/performance/dashboard.py"
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/performance -vs"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,320 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Track performance results and compare against historical baseline.
|
||||
|
||||
This script:
|
||||
1) reads current benchmark results from fastvideo/tests/performance/results,
|
||||
2) writes normalized tracking records to the Modal volume path,
|
||||
3) compares each current record against the mean of up to 5 prior records,
|
||||
4) exits non-zero if any metric regresses by more than 15%.
|
||||
"""
|
||||
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import statistics
|
||||
import sys
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from hf_store import sync_from_hf, upload_record, load_records_for_model, sanitize, safe_float
|
||||
|
||||
# Use the env var passed by Modal, fallback to a default if needed
|
||||
HF_REPO_ID = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
|
||||
HF_TOKEN = os.environ.get("HF_API_KEY")
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
"results",
|
||||
)
|
||||
TRACKING_ROOT = os.environ.get(
|
||||
"PERFORMANCE_TRACKING_ROOT",
|
||||
"/tmp/perf-tracking",
|
||||
)
|
||||
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
# test_scope = os.environ.get("TEST_SCOPE", "")
|
||||
# branch = os.environ.get("BUILDKITE_BRANCH", "")
|
||||
# return test_scope == "full" and branch == "main"
|
||||
return True # only for testing purpose.
|
||||
|
||||
def _sanitize(value: str) -> str:
|
||||
return re.sub(r"[^A-Za-z0-9._-]", "_", value)
|
||||
|
||||
|
||||
def _safe_float(value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _load_current_results() -> list[dict[str, Any]]:
|
||||
pattern = os.path.join(RESULTS_DIR, "perf_*.json")
|
||||
records: list[dict[str, Any]] = []
|
||||
for path in sorted(glob.glob(pattern)):
|
||||
with open(path, encoding="utf-8") as f:
|
||||
records.append(json.load(f))
|
||||
return records
|
||||
|
||||
|
||||
def _normalize_record(result: dict[str, Any]) -> dict[str, Any]:
|
||||
benchmark_id = result.get("benchmark_id", "unknown")
|
||||
model_id = benchmark_id
|
||||
|
||||
timestamp = result.get("timestamp")
|
||||
if not timestamp:
|
||||
timestamp = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
commit_sha = result.get("commit") or os.environ.get("BUILDKITE_COMMIT", "")
|
||||
latency = _safe_float(result.get("avg_generation_time_s"))
|
||||
throughput = _safe_float(result.get("throughput_fps"))
|
||||
memory = _safe_float(result.get("max_peak_memory_mb"))
|
||||
|
||||
return {
|
||||
"model_id": model_id,
|
||||
"timestamp": timestamp,
|
||||
"commit_sha": commit_sha,
|
||||
"gpu_type": result.get("device", "unknown"),
|
||||
"latency": latency,
|
||||
"throughput": throughput,
|
||||
"memory": memory,
|
||||
"success": True,
|
||||
}
|
||||
|
||||
|
||||
def _write_tracking_record(record: dict[str, Any]) -> str:
|
||||
model_dir = os.path.join(TRACKING_ROOT, _sanitize(record["model_id"]))
|
||||
os.makedirs(model_dir, exist_ok=True)
|
||||
|
||||
timestamp = _sanitize(record["timestamp"])
|
||||
commit = _sanitize(record["commit_sha"] or "unknown")
|
||||
out_path = os.path.join(model_dir, f"{timestamp}_{commit}.json")
|
||||
|
||||
with open(out_path, "w", encoding="utf-8") as f:
|
||||
json.dump(record, f, indent=2)
|
||||
|
||||
return out_path
|
||||
|
||||
def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
|
||||
values = [
|
||||
_safe_float(r.get(key))
|
||||
for r in records
|
||||
]
|
||||
values = [v for v in values if v is not None]
|
||||
if not values:
|
||||
return None
|
||||
return statistics.median(values)
|
||||
|
||||
|
||||
def _check_regressions(
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
) -> list[str]:
|
||||
failures: list[str] = []
|
||||
|
||||
for metric in ("latency", "memory"):
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
curr = _safe_float(current.get(metric))
|
||||
if baseline is None or curr is None or baseline <= 0:
|
||||
continue
|
||||
regression = (curr - baseline) / baseline
|
||||
if regression > max_regression:
|
||||
failures.append(
|
||||
f"{current['model_id']} {metric} regressed by {regression * 100:.1f}% "
|
||||
f"(current={curr:.3f}, baseline_median={baseline:.3f})"
|
||||
)
|
||||
|
||||
baseline_tp = _baseline_metric(baseline_records, "throughput")
|
||||
curr_tp = _safe_float(current.get("throughput"))
|
||||
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
|
||||
regression = (baseline_tp - curr_tp) / baseline_tp
|
||||
if regression > max_regression:
|
||||
failures.append(
|
||||
f"{current['model_id']} throughput regressed by {regression * 100:.1f}% "
|
||||
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})"
|
||||
)
|
||||
|
||||
return failures
|
||||
|
||||
|
||||
def _metric_delta_percent(
|
||||
metric: str,
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
) -> float | None:
|
||||
curr = _safe_float(current.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
if curr is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
|
||||
if metric in ("latency", "memory"):
|
||||
return (curr - baseline) / baseline * 100.0
|
||||
if metric == "throughput":
|
||||
return (baseline - curr) / baseline * 100.0
|
||||
return None
|
||||
|
||||
|
||||
def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
if value is None:
|
||||
return "n/a"
|
||||
return f"{value:.{precision}f}"
|
||||
|
||||
def _build_summary_row(
|
||||
record: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
has_failed: bool
|
||||
) -> dict[str, Any]:
|
||||
"""Formats a single benchmark result into a row for the Markdown summary table."""
|
||||
|
||||
latency_base = _safe_float(_baseline_metric(baseline_records, "latency"))
|
||||
throughput_base = _safe_float(_baseline_metric(baseline_records, "throughput"))
|
||||
memory_base = _safe_float(_baseline_metric(baseline_records, "memory"))
|
||||
|
||||
# Calculate percentages for the 'Worst Regression' column
|
||||
latency_reg = _metric_delta_percent("latency", record, baseline_records)
|
||||
throughput_reg = _metric_delta_percent("throughput", record, baseline_records)
|
||||
memory_reg = _metric_delta_percent("memory", record, baseline_records)
|
||||
|
||||
regressions = [v for v in (latency_reg, throughput_reg, memory_reg) if v is not None]
|
||||
worst_regression_pct = max(regressions) if regressions else None
|
||||
|
||||
return {
|
||||
"model_id": record["model_id"],
|
||||
"gpu_type": record["gpu_type"],
|
||||
"baseline_n": len(baseline_records),
|
||||
"latency_curr": _safe_float(record.get("latency")),
|
||||
"latency_base": latency_base,
|
||||
"throughput_curr": _safe_float(record.get("throughput")),
|
||||
"throughput_base": throughput_base,
|
||||
"memory_curr": _safe_float(record.get("memory")),
|
||||
"memory_base": memory_base,
|
||||
"worst_regression_pct": worst_regression_pct,
|
||||
"failed": has_failed,
|
||||
}
|
||||
|
||||
def _build_markdown_summary(
|
||||
summary_rows: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
) -> str:
|
||||
lines = [
|
||||
"## Performance Baseline Comparison",
|
||||
"",
|
||||
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
|
||||
"",
|
||||
"| Model | GPU | Baseline N | Latency (curr/base) | Throughput (curr/base) | Memory (curr/base) | Worst Regression | Status |",
|
||||
"|---|---|---:|---|---|---|---:|---|",
|
||||
]
|
||||
|
||||
for row in summary_rows:
|
||||
latency = f"{_compact_value(row['latency_curr'])} / {_compact_value(row['latency_base'])}"
|
||||
throughput = f"{_compact_value(row['throughput_curr'])} / {_compact_value(row['throughput_base'])}"
|
||||
memory = f"{_compact_value(row['memory_curr'], 1)} / {_compact_value(row['memory_base'], 1)}"
|
||||
|
||||
worst_reg = "n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%"
|
||||
status = "FAIL" if row["failed"] else "PASS"
|
||||
|
||||
lines.append(
|
||||
f"| {row['model_id']} | {row['gpu_type']} | {row['baseline_n']} | "
|
||||
f"{latency} | {throughput} | {memory} | {worst_reg} | {status} |"
|
||||
)
|
||||
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
def _emit_markdown_summary(markdown: str, commit_sha: str) -> None:
|
||||
print("\n" + markdown)
|
||||
|
||||
# 1. Existing GitHub logic (safe to keep)
|
||||
summary_path = os.environ.get("GITHUB_STEP_SUMMARY")
|
||||
if summary_path:
|
||||
with open(summary_path, "a", encoding="utf-8") as f:
|
||||
f.write(markdown + "\n")
|
||||
|
||||
# 2. Write to Modal volume for Buildkite to pick up in post-run hook
|
||||
try:
|
||||
perf_reports_dir = "/root/data/perf_reports"
|
||||
os.makedirs(perf_reports_dir, exist_ok=True)
|
||||
short_sha = commit_sha[:7] if commit_sha else "unknown"
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
report_path = os.path.join(perf_reports_dir, f"perf_{short_sha}_{timestamp}.md")
|
||||
with open(report_path, "w", encoding="utf-8") as f:
|
||||
f.write(markdown + "\n")
|
||||
print(f"Performance report written to {report_path}")
|
||||
except Exception as e:
|
||||
print(f"Failed to write performance report to Modal volume: {e}")
|
||||
|
||||
def main() -> int:
|
||||
# Pull the current state of the world from HF
|
||||
sync_from_hf(TRACKING_ROOT)
|
||||
|
||||
current_results = _load_current_results()
|
||||
if not current_results:
|
||||
print(f"No performance result files found in {RESULTS_DIR}")
|
||||
return 0
|
||||
|
||||
all_failures = []
|
||||
summary_rows = []
|
||||
persist_tracking = _should_persist_tracking()
|
||||
|
||||
if persist_tracking:
|
||||
print("Tracking persistence enabled: full-suite run on main branch")
|
||||
else:
|
||||
print("Tracking persistence disabled: only full-suite runs on main branch are persisted")
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
|
||||
baseline_records = load_records_for_model(
|
||||
TRACKING_ROOT, record["model_id"], record["gpu_type"],
|
||||
last_n=5, successful_only=True
|
||||
)
|
||||
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
|
||||
# Tag the current record based on the failure.
|
||||
if not baseline_records:
|
||||
# INITIALIZATION CASE: First run for this model/GPU
|
||||
print(f"No baseline for {record['model_id']} on {record['gpu_type']}. Initializing...")
|
||||
failures = []
|
||||
record["success"] = True # The first run is always "successful"
|
||||
else:
|
||||
# COMPARISON CASE: Compare against the mean of the last 5 good runs
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
if failures:
|
||||
record["success"] = False
|
||||
all_failures.extend(failures)
|
||||
else:
|
||||
record["success"] = True
|
||||
|
||||
# 5. Persist to HF if we are on main
|
||||
if persist_tracking:
|
||||
# This writes the JSON with the "success" field to /tmp
|
||||
current_path = _write_tracking_record(record)
|
||||
# This pushes it to the FastVideo/performance-tracking repo
|
||||
upload_record(current_path, record)
|
||||
|
||||
|
||||
summary_row = _build_summary_row(record, baseline_records, bool(failures))
|
||||
summary_rows.append(summary_row)
|
||||
|
||||
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
|
||||
|
||||
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
|
||||
_emit_markdown_summary(markdown, commit_sha)
|
||||
|
||||
if all_failures:
|
||||
print("Performance regression check failed:")
|
||||
for item in all_failures:
|
||||
print(f" - {item}")
|
||||
return 1
|
||||
|
||||
print("Performance baseline comparison passed")
|
||||
return 0
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -1,102 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
from datetime import datetime
|
||||
|
||||
import plotly.express as px
|
||||
import pandas as pd
|
||||
|
||||
from hf_store import sync_from_hf, load_as_dataframe
|
||||
|
||||
# -----------------------------
|
||||
# 1. Grouping
|
||||
# -----------------------------
|
||||
def group_data(df: pd.DataFrame):
|
||||
# Group only by model+GPU so each group produces a time-series line.
|
||||
# config_id (commit SHA) is carried as a column for hover/color use.
|
||||
keys = ["model_id", "gpu_type"]
|
||||
return df.groupby(keys, dropna=False)
|
||||
|
||||
# -----------------------------
|
||||
# 2. Plot builder
|
||||
# -----------------------------
|
||||
def build_plots(df: pd.DataFrame) -> list:
|
||||
figs = []
|
||||
|
||||
for (model_id, gpu_type), g in group_data(df):
|
||||
g = g.sort_values("timestamp")
|
||||
|
||||
# One chart per metric so the y-axes aren't on wildly different scales
|
||||
for metric in ("latency", "throughput", "memory"):
|
||||
if g[metric].isna().all():
|
||||
continue
|
||||
|
||||
fig = px.line(
|
||||
g,
|
||||
x="timestamp",
|
||||
y=metric,
|
||||
markers=True,
|
||||
hover_data=["config_id", "commit_sha"],
|
||||
title=f"{model_id} | {gpu_type} | {metric}",
|
||||
labels={"timestamp": "Time", metric: metric},
|
||||
)
|
||||
figs.append(fig)
|
||||
|
||||
return figs
|
||||
|
||||
# -----------------------------
|
||||
# 3. Render HTML dashboard
|
||||
# -----------------------------
|
||||
def render_html(figs: list, days: int) -> str:
|
||||
html_parts = [
|
||||
"<html>",
|
||||
"<head><meta charset='utf-8'>",
|
||||
"<style>body { font-family: sans-serif; margin: 2rem; }</style>",
|
||||
"</head><body>",
|
||||
f"<h2>Performance Dashboard (last {days} days)</h2>",
|
||||
]
|
||||
|
||||
for fig in figs:
|
||||
html_parts.append(fig.to_html(full_html=False, include_plotlyjs="cdn"))
|
||||
|
||||
html_parts.append("</body></html>")
|
||||
return "\n".join(html_parts)
|
||||
|
||||
# -----------------------------
|
||||
# 5. Main
|
||||
# -----------------------------
|
||||
def main() -> None:
|
||||
days = int(os.environ.get("DASHBOARD_DAYS", "30"))
|
||||
|
||||
local_dir = sync_from_hf("/tmp/perf-tracking")
|
||||
df = load_as_dataframe(local_dir, days=days)
|
||||
|
||||
if df.empty:
|
||||
print("No data found")
|
||||
return
|
||||
|
||||
# Sanity-check: log what we actually loaded
|
||||
print(f"Loaded {len(df)} records across {df['model_id'].nunique()} model(s), "
|
||||
f"{df['gpu_type'].nunique()} GPU type(s), "
|
||||
f"date range: {df['timestamp'].min()} → {df['timestamp'].max()}")
|
||||
|
||||
figs = build_plots(df)
|
||||
html = render_html(figs, days)
|
||||
|
||||
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
|
||||
report_dir = "/root/data/perf_reports"
|
||||
os.makedirs(report_dir, exist_ok=True)
|
||||
|
||||
filename = f"dashboard_{commit_sha}_{timestamp}.html"
|
||||
output_file = os.path.join(report_dir, filename)
|
||||
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
f.write(html)
|
||||
|
||||
print(f"Dashboard generated: {output_file}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,254 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared HuggingFace storage utilities for performance tracking.
|
||||
|
||||
Provides a single place for:
|
||||
- Syncing the HF dataset repo to a local directory
|
||||
- Loading raw JSON records (with optional recency filter)
|
||||
- Loading records as a normalized pandas DataFrame
|
||||
- Uploading individual result files back to HF
|
||||
- Common helpers: sanitize, safe_float
|
||||
"""
|
||||
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Configuration — read once at import time, shared across both consumers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
HF_REPO_ID: str = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
|
||||
HF_TOKEN: str | None = os.environ.get("HF_API_KEY")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Low-level helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def sanitize(value: str) -> str:
|
||||
"""Return a filesystem- and HF-path-safe version of *value*."""
|
||||
return re.sub(r"[^A-Za-z0-9._-]", "_", value)
|
||||
|
||||
|
||||
def safe_float(value: Any) -> float | None:
|
||||
"""Coerce *value* to float, returning None on failure."""
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HF I/O
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def sync_from_hf(local_dir: str) -> str:
|
||||
"""Download the HF dataset repo snapshot to *local_dir*.
|
||||
|
||||
Returns *local_dir* so callers can chain: ``load_records(sync_from_hf(...))``.
|
||||
On failure (empty repo, no credentials, network error) the function logs a
|
||||
warning and returns *local_dir* unchanged so the caller can still work with
|
||||
whatever is already on disk.
|
||||
"""
|
||||
if not HF_REPO_ID:
|
||||
print("hf_store: HF_REPO_ID not set, skipping sync.")
|
||||
return local_dir
|
||||
|
||||
print(f"hf_store: syncing from {HF_REPO_ID} → {local_dir}")
|
||||
try:
|
||||
snapshot_download(
|
||||
repo_id=HF_REPO_ID,
|
||||
repo_type="dataset",
|
||||
local_dir=local_dir,
|
||||
token=HF_TOKEN,
|
||||
allow_patterns="*.json",
|
||||
)
|
||||
except Exception as exc:
|
||||
print(f"hf_store: sync skipped — {exc}")
|
||||
|
||||
return local_dir
|
||||
|
||||
|
||||
def upload_record(local_path: str, record: dict[str, Any]) -> None:
|
||||
"""Upload *local_path* to the HF repo under ``<model_id>/<filename>``.
|
||||
|
||||
Silently skips if HF_TOKEN is absent so local/CI runs without credentials
|
||||
don't crash.
|
||||
"""
|
||||
if not HF_TOKEN:
|
||||
print("hf_store: HF_API_KEY not set, skipping upload.")
|
||||
return
|
||||
|
||||
model_id = record.get("model_id", "unknown")
|
||||
path_in_repo = f"{sanitize(model_id)}/{os.path.basename(local_path)}"
|
||||
commit_sha = (record.get("commit_sha") or "unknown")[:7]
|
||||
|
||||
api = HfApi(token=HF_TOKEN)
|
||||
try:
|
||||
api.upload_file(
|
||||
path_or_fileobj=local_path,
|
||||
path_in_repo=path_in_repo,
|
||||
repo_id=HF_REPO_ID,
|
||||
repo_type="dataset",
|
||||
commit_message=f"Perf: {model_id} at {commit_sha}",
|
||||
)
|
||||
print(f"hf_store: uploaded → {HF_REPO_ID}/{path_in_repo}")
|
||||
except Exception as exc:
|
||||
print(f"hf_store: upload failed — {exc}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Record loading
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def load_records(
|
||||
local_dir: str,
|
||||
*,
|
||||
days: int | None = None,
|
||||
successful_only: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return raw JSON dicts from *local_dir*.
|
||||
|
||||
Args:
|
||||
local_dir: Root directory previously populated by :func:`sync_from_hf`.
|
||||
days: When set, discard records whose ``timestamp`` is older than this
|
||||
many days. Records with a missing/unparseable timestamp are kept.
|
||||
successful_only: When True, only records with ``success=True`` are
|
||||
returned. Useful when building a regression baseline.
|
||||
|
||||
Returns:
|
||||
List of raw dicts sorted by ``timestamp`` ascending (records that could
|
||||
not be parsed are silently skipped).
|
||||
"""
|
||||
cutoff: datetime | None = None
|
||||
if days is not None:
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
|
||||
|
||||
records: list[dict[str, Any]] = []
|
||||
|
||||
for path in sorted(glob.glob(os.path.join(local_dir, "**", "*.json"), recursive=True)):
|
||||
try:
|
||||
with open(path, encoding="utf-8") as fh:
|
||||
data: dict[str, Any] = json.load(fh)
|
||||
except (OSError, json.JSONDecodeError):
|
||||
continue
|
||||
|
||||
if successful_only and not data.get("success", True):
|
||||
continue
|
||||
|
||||
if cutoff is not None:
|
||||
raw_ts = data.get("timestamp")
|
||||
if raw_ts:
|
||||
try:
|
||||
ts = datetime.fromisoformat(raw_ts)
|
||||
if ts.tzinfo is None:
|
||||
ts = ts.replace(tzinfo=timezone.utc)
|
||||
if ts < cutoff:
|
||||
continue
|
||||
except ValueError:
|
||||
pass # keep records with unparseable timestamps
|
||||
|
||||
records.append(data)
|
||||
|
||||
return records
|
||||
|
||||
|
||||
def load_records_for_model(
|
||||
local_dir: str,
|
||||
model_id: str,
|
||||
gpu_type: str | None = None,
|
||||
*,
|
||||
last_n: int | None = None,
|
||||
successful_only: bool = True,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return records for a specific *model_id*, optionally filtered by GPU.
|
||||
|
||||
Args:
|
||||
local_dir: Root directory previously populated by :func:`sync_from_hf`.
|
||||
model_id: Matches the ``model_id`` field inside each JSON record.
|
||||
gpu_type: When set, only records whose ``gpu_type`` matches are returned.
|
||||
last_n: When set, return only the most recent *n* records (after all
|
||||
other filters). Useful for sliding-window baseline calculations.
|
||||
successful_only: Passed through to :func:`load_records`.
|
||||
|
||||
Returns:
|
||||
List of matching dicts sorted by timestamp ascending.
|
||||
"""
|
||||
model_dir = os.path.join(local_dir, sanitize(model_id))
|
||||
if not os.path.isdir(model_dir):
|
||||
return []
|
||||
|
||||
records = load_records(model_dir, successful_only=successful_only)
|
||||
|
||||
if gpu_type is not None:
|
||||
records = [r for r in records if r.get("gpu_type") == gpu_type]
|
||||
|
||||
if last_n is not None:
|
||||
records = records[-last_n:]
|
||||
|
||||
return records
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DataFrame helpers (dashboard / analytics consumers)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_NUMERIC_COLS = ("latency", "throughput", "memory")
|
||||
|
||||
|
||||
def normalize_dataframe(df: pd.DataFrame) -> pd.DataFrame:
|
||||
"""Apply standard type coercions to a raw records DataFrame.
|
||||
|
||||
- Parses ``timestamp`` to UTC-aware datetime.
|
||||
- Coerces ``latency``, ``throughput``, ``memory`` to float.
|
||||
- Adds a ``config_id`` column (first 7 chars of ``commit_sha``).
|
||||
|
||||
Returns the mutated DataFrame (also modifies in place for efficiency).
|
||||
"""
|
||||
if df.empty:
|
||||
return df
|
||||
|
||||
df["timestamp"] = pd.to_datetime(df["timestamp"], utc=True, errors="coerce")
|
||||
df["config_id"] = df.get("commit_sha", pd.Series(dtype=str)).fillna("unknown").str[:7]
|
||||
|
||||
for col in _NUMERIC_COLS:
|
||||
if col in df.columns:
|
||||
df[col] = pd.to_numeric(df[col], errors="coerce")
|
||||
|
||||
return df
|
||||
|
||||
|
||||
def load_as_dataframe(
|
||||
local_dir: str,
|
||||
*,
|
||||
days: int | None = None,
|
||||
successful_only: bool = False,
|
||||
) -> pd.DataFrame:
|
||||
"""Load and normalize records from *local_dir* into a pandas DataFrame.
|
||||
|
||||
Combines :func:`load_records` + :func:`normalize_dataframe` into a single
|
||||
call for consumers (e.g. the dashboard) that work exclusively with
|
||||
DataFrames.
|
||||
|
||||
Args:
|
||||
local_dir: Root directory previously populated by :func:`sync_from_hf`.
|
||||
days: Passed through to :func:`load_records`.
|
||||
successful_only: Passed through to :func:`load_records`.
|
||||
|
||||
Returns:
|
||||
Normalized DataFrame, or an empty DataFrame if no records were found.
|
||||
"""
|
||||
records = load_records(local_dir, days=days, successful_only=successful_only)
|
||||
if not records:
|
||||
return pd.DataFrame()
|
||||
|
||||
df = pd.DataFrame(records)
|
||||
return normalize_dataframe(df)
|
||||
@@ -164,10 +164,6 @@ def test_inference_performance(cfg):
|
||||
avg_time = sum(times) / len(times)
|
||||
max_peak_memory = max(peak_memories)
|
||||
device_name = torch.cuda.get_device_name()
|
||||
num_frames = gen_kwargs.get("num_frames")
|
||||
throughput_fps = (1.0 / avg_time) if avg_time > 0 else None
|
||||
if isinstance(num_frames, (int, float)) and avg_time > 0:
|
||||
throughput_fps = num_frames / avg_time
|
||||
|
||||
results = {
|
||||
"benchmark_id": cfg["benchmark_id"],
|
||||
@@ -178,8 +174,6 @@ def test_inference_performance(cfg):
|
||||
"num_measurement_runs": num_measure,
|
||||
"avg_generation_time_s": round(avg_time, 3),
|
||||
"individual_times_s": [round(t, 3) for t in times],
|
||||
"throughput_fps": round(throughput_fps, 3)
|
||||
if throughput_fps is not None else None,
|
||||
"max_peak_memory_mb": round(max_peak_memory, 1),
|
||||
"individual_peak_memories_mb": [round(m, 1) for m in peak_memories],
|
||||
"thresholds": thresholds,
|
||||
|
||||
@@ -10,21 +10,21 @@ Note: num_inference_steps is reduced to 4 for faster CI.
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
build_generated_output_dir,
|
||||
build_reference_folder_path,
|
||||
get_cuda_device_name,
|
||||
resolve_device_reference_folder,
|
||||
select_ssim_params,
|
||||
)
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
build_generated_output_dir,
|
||||
build_reference_folder_path,
|
||||
get_cuda_device_name,
|
||||
resolve_device_reference_folder,
|
||||
select_ssim_params,
|
||||
)
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -48,20 +48,20 @@ def _find_lingbotworld_examples_root() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
device_name = get_cuda_device_name()
|
||||
device_reference_folder = resolve_device_reference_folder(
|
||||
(
|
||||
("A40", "A40"),
|
||||
("L40S", "L40S"),
|
||||
("H100", "H100"),
|
||||
("H200", "H200"),
|
||||
),
|
||||
device_name=device_name,
|
||||
logger=logger,
|
||||
)
|
||||
device_name = get_cuda_device_name()
|
||||
device_reference_folder = resolve_device_reference_folder(
|
||||
(
|
||||
("A40", "A40"),
|
||||
("L40S", "L40S"),
|
||||
("H100", "H100"),
|
||||
("H200", "H200"),
|
||||
),
|
||||
device_name=device_name,
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
|
||||
LINGBOT_PARAMS = {
|
||||
LINGBOT_PARAMS = {
|
||||
"model_path": "FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
"num_gpus": 2,
|
||||
"height": 256,
|
||||
@@ -88,29 +88,29 @@ LINGBOT_PARAMS = {
|
||||
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
|
||||
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
|
||||
"皮肤,肢体,面部特征,汽车,电线"
|
||||
),
|
||||
}
|
||||
_LINGBOT_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
|
||||
LINGBOT_PARAMS["model_path"])
|
||||
LINGBOT_FULL_QUALITY_PARAMS = {
|
||||
"model_path": LINGBOT_PARAMS["model_path"],
|
||||
"num_gpus": LINGBOT_PARAMS["num_gpus"],
|
||||
"height": _LINGBOT_FULL_QUALITY_DEFAULTS.height,
|
||||
"width": _LINGBOT_FULL_QUALITY_DEFAULTS.width,
|
||||
"num_frames": LINGBOT_PARAMS["num_frames"], # default num_frames: 125
|
||||
"num_inference_steps": _LINGBOT_FULL_QUALITY_DEFAULTS.num_inference_steps,
|
||||
"guidance_scale": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale,
|
||||
"guidance_scale_2": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale_2,
|
||||
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
|
||||
"flow_shift": LINGBOT_PARAMS["flow_shift"],
|
||||
"boundary_ratio": _LINGBOT_FULL_QUALITY_DEFAULTS.boundary_ratio,
|
||||
"seed": _LINGBOT_FULL_QUALITY_DEFAULTS.seed,
|
||||
"fps": _LINGBOT_FULL_QUALITY_DEFAULTS.fps,
|
||||
"spatial_scale": LINGBOT_PARAMS["spatial_scale"],
|
||||
"example_case": LINGBOT_PARAMS["example_case"],
|
||||
"image_path": LINGBOT_PARAMS["image_path"],
|
||||
"negative_prompt": _LINGBOT_FULL_QUALITY_DEFAULTS.negative_prompt,
|
||||
}
|
||||
),
|
||||
}
|
||||
_LINGBOT_FULL_QUALITY_DEFAULTS = SamplingParam.from_pretrained(
|
||||
LINGBOT_PARAMS["model_path"])
|
||||
LINGBOT_FULL_QUALITY_PARAMS = {
|
||||
"model_path": LINGBOT_PARAMS["model_path"],
|
||||
"num_gpus": LINGBOT_PARAMS["num_gpus"],
|
||||
"height": _LINGBOT_FULL_QUALITY_DEFAULTS.height,
|
||||
"width": _LINGBOT_FULL_QUALITY_DEFAULTS.width,
|
||||
"num_frames": LINGBOT_PARAMS["num_frames"], # default num_frames: 125
|
||||
"num_inference_steps": _LINGBOT_FULL_QUALITY_DEFAULTS.num_inference_steps,
|
||||
"guidance_scale": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale,
|
||||
"guidance_scale_2": _LINGBOT_FULL_QUALITY_DEFAULTS.guidance_scale_2,
|
||||
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
|
||||
"flow_shift": LINGBOT_PARAMS["flow_shift"],
|
||||
"boundary_ratio": _LINGBOT_FULL_QUALITY_DEFAULTS.boundary_ratio,
|
||||
"seed": _LINGBOT_FULL_QUALITY_DEFAULTS.seed,
|
||||
"fps": _LINGBOT_FULL_QUALITY_DEFAULTS.fps,
|
||||
"spatial_scale": LINGBOT_PARAMS["spatial_scale"],
|
||||
"example_case": LINGBOT_PARAMS["example_case"],
|
||||
"image_path": LINGBOT_PARAMS["image_path"],
|
||||
"negative_prompt": _LINGBOT_FULL_QUALITY_DEFAULTS.negative_prompt,
|
||||
}
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"The video presents a soaring journey through a fantasy jungle. The wind "
|
||||
@@ -123,80 +123,80 @@ TEST_PROMPTS = [
|
||||
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
params = select_ssim_params(LINGBOT_PARAMS, LINGBOT_FULL_QUALITY_PARAMS)
|
||||
|
||||
if device_reference_folder is None:
|
||||
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
|
||||
if torch.cuda.device_count() < params["num_gpus"]:
|
||||
pytest.skip(
|
||||
f"LingBot SSIM test requires {params['num_gpus']} GPUs, "
|
||||
f"but only {torch.cuda.device_count()} detected."
|
||||
)
|
||||
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
params = select_ssim_params(LINGBOT_PARAMS, LINGBOT_FULL_QUALITY_PARAMS)
|
||||
|
||||
if device_reference_folder is None:
|
||||
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
|
||||
if torch.cuda.device_count() < params["num_gpus"]:
|
||||
pytest.skip(
|
||||
f"LingBot SSIM test requires {params['num_gpus']} GPUs, "
|
||||
f"but only {torch.cuda.device_count()} detected."
|
||||
)
|
||||
|
||||
examples_root = _find_lingbotworld_examples_root()
|
||||
if examples_root is None:
|
||||
pytest.skip(
|
||||
"lingbotworld_examples not found under examples/inference/basic.")
|
||||
|
||||
action_path = os.path.join(examples_root, params["example_case"])
|
||||
action_path = os.path.join(examples_root, params["example_case"])
|
||||
if not (os.path.exists(os.path.join(action_path, "poses.npy"))
|
||||
and os.path.exists(os.path.join(action_path, "intrinsics.npy"))):
|
||||
pytest.skip(f"Missing camera npy files under {action_path}")
|
||||
|
||||
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
|
||||
action_path=action_path,
|
||||
num_frames=params["num_frames"],
|
||||
height=params["height"],
|
||||
width=params["width"],
|
||||
spatial_scale=params["spatial_scale"],
|
||||
)
|
||||
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
|
||||
action_path=action_path,
|
||||
num_frames=params["num_frames"],
|
||||
height=params["height"],
|
||||
width=params["width"],
|
||||
spatial_scale=params["spatial_scale"],
|
||||
)
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
model_id = "LingBot-World-Base-Cam-Diffusers"
|
||||
output_dir = build_generated_output_dir(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
ATTENTION_BACKEND,
|
||||
)
|
||||
output_dir = build_generated_output_dir(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
ATTENTION_BACKEND,
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": params["num_gpus"],
|
||||
"flow_shift": params["flow_shift"],
|
||||
"boundary_ratio": params["boundary_ratio"],
|
||||
"use_fsdp_inference": True,
|
||||
"dit_cpu_offload": True,
|
||||
init_kwargs = {
|
||||
"num_gpus": params["num_gpus"],
|
||||
"flow_shift": params["flow_shift"],
|
||||
"boundary_ratio": params["boundary_ratio"],
|
||||
"use_fsdp_inference": True,
|
||||
"dit_cpu_offload": True,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"vae_cpu_offload": False,
|
||||
"pin_cpu_memory": True,
|
||||
}
|
||||
generation_kwargs = {
|
||||
"output_path": output_dir,
|
||||
"image_path": params["image_path"],
|
||||
"height": params["height"],
|
||||
"width": params["width"],
|
||||
"num_frames": aligned_num_frames,
|
||||
"num_inference_steps": params["num_inference_steps"],
|
||||
"guidance_scale": params["guidance_scale"],
|
||||
"guidance_scale_2": params["guidance_scale_2"],
|
||||
"embedded_cfg_scale": params["embedded_cfg_scale"],
|
||||
"seed": params["seed"],
|
||||
"fps": params["fps"],
|
||||
"negative_prompt": params["negative_prompt"],
|
||||
"c2ws_plucker_emb": c2ws_plucker_emb,
|
||||
}
|
||||
generation_kwargs = {
|
||||
"output_path": output_dir,
|
||||
"image_path": params["image_path"],
|
||||
"height": params["height"],
|
||||
"width": params["width"],
|
||||
"num_frames": aligned_num_frames,
|
||||
"num_inference_steps": params["num_inference_steps"],
|
||||
"guidance_scale": params["guidance_scale"],
|
||||
"guidance_scale_2": params["guidance_scale_2"],
|
||||
"embedded_cfg_scale": params["embedded_cfg_scale"],
|
||||
"seed": params["seed"],
|
||||
"fps": params["fps"],
|
||||
"negative_prompt": params["negative_prompt"],
|
||||
"c2ws_plucker_emb": c2ws_plucker_emb,
|
||||
}
|
||||
|
||||
generator: VideoGenerator | None = None
|
||||
try:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=params["model_path"], **init_kwargs)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=params["model_path"], **init_kwargs)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
finally:
|
||||
if generator is not None:
|
||||
generator.shutdown()
|
||||
@@ -205,12 +205,12 @@ def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
|
||||
assert os.path.exists(generated_video_path), (
|
||||
f"Output video was not generated at {generated_video_path}")
|
||||
|
||||
reference_folder = build_reference_folder_path(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
ATTENTION_BACKEND,
|
||||
)
|
||||
reference_folder = build_reference_folder_path(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
ATTENTION_BACKEND,
|
||||
)
|
||||
if not os.path.exists(reference_folder):
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}")
|
||||
@@ -234,11 +234,11 @@ def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info("SSIM mean value: %s", mean_ssim)
|
||||
|
||||
write_ssim_results(output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path,
|
||||
params["num_inference_steps"], prompt)
|
||||
write_ssim_results(output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path,
|
||||
params["num_inference_steps"], prompt)
|
||||
|
||||
min_acceptable_ssim = 0.70
|
||||
min_acceptable_ssim = 0.90
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
|
||||
f"for {model_id} with backend {ATTENTION_BACKEND}")
|
||||
|
||||
@@ -89,5 +89,5 @@ def test_ltx2_distilled_inference_similarity(
|
||||
model_id=model_id,
|
||||
default_params_map=LTX2_DISTILLED_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_LTX2_DISTILLED_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.60,
|
||||
min_acceptable_ssim=0.98,
|
||||
)
|
||||
|
||||
@@ -145,7 +145,6 @@ TURBODIFFUSION_I2V_IMAGE_PATHS = [
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="Disabled: causes OOM too often in CI")
|
||||
@pytest.mark.parametrize("prompt", TURBODIFFUSION_I2V_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize(
|
||||
"model_id",
|
||||
|
||||
@@ -1,91 +0,0 @@
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.pipelines.stages.image_encoding import ImageVAEEncodingStage
|
||||
|
||||
|
||||
def make_stage() -> ImageVAEEncodingStage:
|
||||
# Bypass __init__: preprocess() does not use self.vae.
|
||||
return ImageVAEEncodingStage.__new__(ImageVAEEncodingStage)
|
||||
|
||||
|
||||
def test_preprocess_pil_image():
|
||||
stage = make_stage()
|
||||
arr = np.array(
|
||||
[[[0, 0, 0], [128, 128, 128], [255, 255, 255]]],
|
||||
dtype=np.uint8,
|
||||
)
|
||||
image = PIL.Image.fromarray(arr, mode="RGB")
|
||||
|
||||
out = stage.preprocess(image, vae_scale_factor=1, height=1, width=3)
|
||||
|
||||
assert out.dtype == torch.float32
|
||||
assert out.shape == (1, 3, 1, 3)
|
||||
torch.testing.assert_close(
|
||||
out[0, 0, 0],
|
||||
torch.tensor([-1.0, 128.0 / 255.0 * 2 - 1, 1.0]),
|
||||
atol=1e-6,
|
||||
rtol=0,
|
||||
)
|
||||
|
||||
|
||||
def test_preprocess_uint8_tensor():
|
||||
stage = make_stage()
|
||||
image = torch.tensor(
|
||||
[[[[0, 128, 255]], [[0, 128, 255]], [[0, 128, 255]]]],
|
||||
dtype=torch.uint8,
|
||||
)
|
||||
|
||||
out = stage.preprocess(image, vae_scale_factor=1, height=1, width=3)
|
||||
|
||||
assert out.dtype == torch.float32
|
||||
expected = torch.tensor([-1.0, 128.0 / 255.0 * 2 - 1, 1.0])
|
||||
torch.testing.assert_close(out[0, 0, 0], expected, atol=1e-6, rtol=0)
|
||||
assert out.max().item() <= 1.0
|
||||
assert out.min().item() >= -1.0
|
||||
|
||||
|
||||
def test_preprocess_float01_tensor_matches_uint8_path():
|
||||
stage = make_stage()
|
||||
uint8_image = torch.tensor(
|
||||
[[[[0, 128, 255]], [[0, 128, 255]], [[0, 128, 255]]]],
|
||||
dtype=torch.uint8,
|
||||
)
|
||||
float_image = uint8_image.float() / 255.0
|
||||
|
||||
out_uint8 = stage.preprocess(uint8_image, vae_scale_factor=1, height=1, width=3)
|
||||
out_float = stage.preprocess(float_image, vae_scale_factor=1, height=1, width=3)
|
||||
|
||||
torch.testing.assert_close(out_uint8, out_float, atol=1e-6, rtol=0)
|
||||
|
||||
|
||||
def test_preprocess_already_normalized_passthrough():
|
||||
stage = make_stage()
|
||||
# Already in [-1, 1]; do_normalize branch must be skipped.
|
||||
image = torch.tensor(
|
||||
[[[[-1.0, 0.0, 1.0]], [[-1.0, 0.0, 1.0]], [[-1.0, 0.0, 1.0]]]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
out = stage.preprocess(image, vae_scale_factor=1, height=1, width=3)
|
||||
|
||||
torch.testing.assert_close(out, image, atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_input, expected_exc",
|
||||
[
|
||||
# Float tensor outside [-1, 1] / [0, 1].
|
||||
(torch.tensor([[[[0.0, 1.5]]]], dtype=torch.float32), ValueError),
|
||||
# Non-floating, non-uint8 tensor.
|
||||
(torch.tensor([[[[0, 1]]]], dtype=torch.int32), ValueError),
|
||||
# Wrong outer type.
|
||||
(np.zeros((1, 3, 1, 3), dtype=np.float32), TypeError),
|
||||
],
|
||||
)
|
||||
def test_preprocess_rejects_invalid_inputs(bad_input, expected_exc):
|
||||
stage = make_stage()
|
||||
with pytest.raises(expected_exc):
|
||||
stage.preprocess(bad_input, vae_scale_factor=1, height=1, width=2)
|
||||
@@ -1 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -1,290 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.callback`.
|
||||
|
||||
Covers the ``Callback`` base class no-op contract and the
|
||||
``CallbackDict`` instantiation / dispatch / state-dict logic.
|
||||
|
||||
The concrete callback subclasses (``GradNormClipCallback``,
|
||||
``EMACallback``, ``ValidationCallback``) have their own test files.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.train.callbacks.callback import (
|
||||
Callback,
|
||||
CallbackDict,
|
||||
_BUILTIN_CALLBACKS,
|
||||
)
|
||||
|
||||
|
||||
# ``fastvideo.logger.init_logger`` sets ``propagate=False`` on its
|
||||
# loggers, so the standard ``caplog`` fixture cannot observe them.
|
||||
# This helper attaches a temporary handler directly to the target
|
||||
# logger and yields the captured records.
|
||||
@contextmanager
|
||||
def _capture_logger(
|
||||
name: str, level: int = logging.WARNING
|
||||
) -> Iterator[list[logging.LogRecord]]:
|
||||
logger = logging.getLogger(name)
|
||||
records: list[logging.LogRecord] = []
|
||||
|
||||
class _Handler(logging.Handler):
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
records.append(record)
|
||||
|
||||
handler = _Handler(level=level)
|
||||
prev_level = logger.level
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(level)
|
||||
try:
|
||||
yield records
|
||||
finally:
|
||||
logger.removeHandler(handler)
|
||||
logger.setLevel(prev_level)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _RecordingCallback(Callback):
|
||||
"""Callback that records every hook call into a shared list."""
|
||||
|
||||
def __init__(self, *, tag: str, sink: list[str]) -> None:
|
||||
self._tag = tag
|
||||
self._sink = sink
|
||||
|
||||
def on_train_start(self, method, iteration: int = 0) -> None:
|
||||
self._sink.append(f"{self._tag}:on_train_start:{iteration}")
|
||||
|
||||
def on_training_step_end(
|
||||
self, method, loss_dict, iteration: int = 0
|
||||
) -> None:
|
||||
self._sink.append(f"{self._tag}:on_training_step_end:{iteration}")
|
||||
|
||||
def on_validation_begin(self, method, iteration: int = 0) -> None:
|
||||
self._sink.append(f"{self._tag}:on_validation_begin:{iteration}")
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return {"tag": self._tag, "marker": 7}
|
||||
|
||||
def load_state_dict(self, sd: dict[str, Any]) -> None:
|
||||
self._sink.append(f"{self._tag}:load:{sd.get('marker')}")
|
||||
|
||||
|
||||
class _NotACallback:
|
||||
"""Plain class used to exercise the non-Callback type guard."""
|
||||
|
||||
def __init__(self, **_: Any) -> None:
|
||||
pass
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A. Callback base class
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCallbackBase:
|
||||
|
||||
def test_default_hooks_return_none(self) -> None:
|
||||
cb = Callback()
|
||||
assert cb.on_train_start(method=None) is None
|
||||
assert (
|
||||
cb.on_training_step_end(method=None, loss_dict={}) is None
|
||||
)
|
||||
assert cb.on_before_optimizer_step(method=None) is None
|
||||
assert cb.on_validation_begin(method=None) is None
|
||||
assert cb.on_validation_end(method=None) is None
|
||||
assert cb.on_train_end(method=None) is None
|
||||
|
||||
def test_default_state_dict_round_trip(self) -> None:
|
||||
cb = Callback()
|
||||
assert cb.state_dict() == {}
|
||||
# Default load_state_dict accepts arbitrary state without raising.
|
||||
assert cb.load_state_dict({"unrelated": 1}) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# B. CallbackDict construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCallbackDictInit:
|
||||
|
||||
def test_empty_config(self) -> None:
|
||||
cb_dict = CallbackDict({}, training_config=object())
|
||||
assert cb_dict._callbacks == {}
|
||||
|
||||
def test_builtin_name_resolves_without_target(self) -> None:
|
||||
# ``grad_clip`` is a registered builtin.
|
||||
cfg = {"grad_clip": {"max_grad_norm": 0.5}}
|
||||
tc = object()
|
||||
cb_dict = CallbackDict(cfg, training_config=tc)
|
||||
|
||||
assert "grad_clip" in cb_dict._callbacks
|
||||
from fastvideo.train.callbacks.grad_clip import (
|
||||
GradNormClipCallback,
|
||||
)
|
||||
cb = cb_dict._callbacks["grad_clip"]
|
||||
assert isinstance(cb, GradNormClipCallback)
|
||||
# CallbackDict wires up training_config + back-pointer.
|
||||
assert cb.training_config is tc
|
||||
assert cb._callback_dict is cb_dict
|
||||
|
||||
def test_explicit_target_overrides_name_lookup(self) -> None:
|
||||
cfg = {
|
||||
"anything_goes": {
|
||||
"_target_": (
|
||||
"fastvideo.train.callbacks.grad_clip."
|
||||
"GradNormClipCallback"
|
||||
),
|
||||
"max_grad_norm": 1.0,
|
||||
}
|
||||
}
|
||||
cb_dict = CallbackDict(cfg, training_config=object())
|
||||
assert "anything_goes" in cb_dict._callbacks
|
||||
|
||||
def test_unknown_name_without_target_is_skipped(self) -> None:
|
||||
cfg = {"mystery": {"some_arg": 1}}
|
||||
with _capture_logger(
|
||||
"fastvideo.train.callbacks.callback"
|
||||
) as records:
|
||||
cb_dict = CallbackDict(cfg, training_config=object())
|
||||
assert cb_dict._callbacks == {}
|
||||
assert any(
|
||||
"missing" in r.getMessage() and "mystery" in r.getMessage()
|
||||
for r in records
|
||||
)
|
||||
|
||||
def test_non_callback_target_raises(self) -> None:
|
||||
cfg = {
|
||||
"bad": {
|
||||
"_target_": (
|
||||
"fastvideo.tests.train.callbacks.test_callback."
|
||||
"_NotACallback"
|
||||
)
|
||||
}
|
||||
}
|
||||
with pytest.raises(TypeError, match="expected a Callback"):
|
||||
CallbackDict(cfg, training_config=object())
|
||||
|
||||
def test_builtin_registry_has_expected_entries(self) -> None:
|
||||
# Sanity: protect the builtin registry from silent shrinkage.
|
||||
assert set(_BUILTIN_CALLBACKS) >= {
|
||||
"grad_clip",
|
||||
"validation",
|
||||
"ema",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# C. Dispatch via __getattr__
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCallbackDictDispatch:
|
||||
|
||||
def _build(
|
||||
self, sink: list[str]
|
||||
) -> CallbackDict:
|
||||
cb_dict = CallbackDict({}, training_config=object())
|
||||
cb_dict._callbacks["first"] = _RecordingCallback(
|
||||
tag="first", sink=sink
|
||||
)
|
||||
cb_dict._callbacks["second"] = _RecordingCallback(
|
||||
tag="second", sink=sink
|
||||
)
|
||||
return cb_dict
|
||||
|
||||
def test_dispatch_calls_all_in_insertion_order(self) -> None:
|
||||
sink: list[str] = []
|
||||
cb_dict = self._build(sink)
|
||||
|
||||
cb_dict.on_train_start(method=None, iteration=3)
|
||||
assert sink == [
|
||||
"first:on_train_start:3",
|
||||
"second:on_train_start:3",
|
||||
]
|
||||
|
||||
def test_dispatch_to_hook_some_callbacks_skip(self) -> None:
|
||||
sink: list[str] = []
|
||||
cb_dict = self._build(sink)
|
||||
# The base Callback subclass below only implements one hook;
|
||||
# dispatch should still fan out without raising.
|
||||
|
||||
class _OnlyValidation(Callback):
|
||||
|
||||
def on_validation_end(
|
||||
self, method, iteration: int = 0
|
||||
) -> None:
|
||||
sink.append(f"vend:{iteration}")
|
||||
|
||||
cb_dict._callbacks["only_v"] = _OnlyValidation()
|
||||
cb_dict.on_validation_end(method=None, iteration=11)
|
||||
assert "vend:11" in sink
|
||||
|
||||
def test_dispatch_unknown_hook_is_noop(self) -> None:
|
||||
# Methods that don't exist on any callback should not raise.
|
||||
cb_dict = self._build([])
|
||||
cb_dict.totally_made_up_hook(method=None, iteration=0)
|
||||
|
||||
def test_underscore_attribute_raises(self) -> None:
|
||||
cb_dict = self._build([])
|
||||
with pytest.raises(AttributeError):
|
||||
getattr(cb_dict, "_does_not_exist")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# D. state_dict / load_state_dict
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCallbackDictStateDict:
|
||||
|
||||
def _build(self) -> tuple[CallbackDict, list[str]]:
|
||||
sink: list[str] = []
|
||||
cb_dict = CallbackDict({}, training_config=object())
|
||||
cb_dict._callbacks["first"] = _RecordingCallback(
|
||||
tag="first", sink=sink
|
||||
)
|
||||
cb_dict._callbacks["second"] = _RecordingCallback(
|
||||
tag="second", sink=sink
|
||||
)
|
||||
return cb_dict, sink
|
||||
|
||||
def test_state_dict_returns_per_callback_dict(self) -> None:
|
||||
cb_dict, _ = self._build()
|
||||
state = cb_dict.state_dict()
|
||||
assert set(state) == {"first", "second"}
|
||||
assert state["first"] == {"tag": "first", "marker": 7}
|
||||
assert state["second"] == {"tag": "second", "marker": 7}
|
||||
|
||||
def test_load_state_dict_dispatches_to_each(self) -> None:
|
||||
cb_dict, sink = self._build()
|
||||
cb_dict.load_state_dict(
|
||||
{
|
||||
"first": {"marker": 1},
|
||||
"second": {"marker": 2},
|
||||
}
|
||||
)
|
||||
assert sink == ["first:load:1", "second:load:2"]
|
||||
|
||||
def test_load_state_dict_missing_key_warns_no_raise(self) -> None:
|
||||
cb_dict, sink = self._build()
|
||||
with _capture_logger(
|
||||
"fastvideo.train.callbacks.callback"
|
||||
) as records:
|
||||
cb_dict.load_state_dict({"first": {"marker": 99}})
|
||||
assert sink == ["first:load:99"]
|
||||
assert any(
|
||||
"second" in r.getMessage() and "not found" in r.getMessage()
|
||||
for r in records
|
||||
)
|
||||
@@ -1,254 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.ema`.
|
||||
|
||||
Exercises the EMA lifecycle (lazy init, ``start_iter`` gating, decay
|
||||
math, ``ema_context`` swap, state-dict round-trip) on a tiny CPU
|
||||
``nn.Linear``. ``EMA_FSDP`` works without ``dist.init_process_group``
|
||||
because ``dist.is_initialized()`` returns False and ``_to_local_tensor``
|
||||
falls through to raw tensors for non-DTensor inputs.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.callbacks.ema import EMACallback
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _Student:
|
||||
|
||||
def __init__(self, transformer: torch.nn.Module | None) -> None:
|
||||
self.transformer = transformer
|
||||
|
||||
|
||||
class _RecordingTracker:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.entries: list[tuple[dict[str, Any], int]] = []
|
||||
|
||||
def log(self, payload: dict[str, Any], step: int) -> None:
|
||||
self.entries.append((payload, step))
|
||||
|
||||
|
||||
class _Method:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer: torch.nn.Module | None,
|
||||
tracker: Any | None = None,
|
||||
) -> None:
|
||||
self.student = _Student(transformer)
|
||||
self.tracker = tracker
|
||||
|
||||
|
||||
def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module:
|
||||
m = torch.nn.Linear(4, 2, bias=False)
|
||||
with torch.no_grad():
|
||||
m.weight.fill_(fill)
|
||||
return m
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A. on_train_start
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOnTrainStart:
|
||||
|
||||
def test_initializes_ema_from_student(self) -> None:
|
||||
transformer = _tiny_transformer(fill=0.5)
|
||||
cb = EMACallback(decay=0.9, start_iter=0)
|
||||
cb.on_train_start(_Method(transformer), iteration=0)
|
||||
|
||||
assert cb.student_ema is not None
|
||||
# Shadow shape matches transformer parameter.
|
||||
shadow = cb.student_ema.shadow["weight"]
|
||||
assert shadow.shape == transformer.weight.shape
|
||||
assert torch.allclose(shadow, transformer.weight.detach().cpu())
|
||||
|
||||
def test_missing_transformer_raises(self) -> None:
|
||||
cb = EMACallback()
|
||||
with pytest.raises(ValueError, match="No student transformer"):
|
||||
cb.on_train_start(_Method(transformer=None), iteration=0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# B. on_training_step_end (decay math + start_iter gating)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOnTrainingStepEnd:
|
||||
|
||||
def test_no_op_before_train_start(self) -> None:
|
||||
cb = EMACallback()
|
||||
# student_ema is None until on_train_start.
|
||||
cb.on_training_step_end(
|
||||
_Method(transformer=None), loss_dict={}, iteration=0
|
||||
)
|
||||
assert not cb._ema_started
|
||||
|
||||
def test_skipped_until_start_iter(self) -> None:
|
||||
transformer = _tiny_transformer(fill=1.0)
|
||||
cb = EMACallback(decay=0.5, start_iter=10)
|
||||
cb.on_train_start(_Method(transformer), iteration=0)
|
||||
|
||||
# Mutate transformer to drift it away from initial shadow.
|
||||
with torch.no_grad():
|
||||
transformer.weight.fill_(7.0)
|
||||
|
||||
cb.on_training_step_end(
|
||||
_Method(transformer), loss_dict={}, iteration=5
|
||||
)
|
||||
# Below start_iter: shadow is untouched, _ema_started False.
|
||||
assert not cb._ema_started
|
||||
assert torch.allclose(
|
||||
cb.student_ema.shadow["weight"],
|
||||
torch.full((2, 4), 1.0),
|
||||
)
|
||||
|
||||
def test_first_active_step_reinits_then_updates(self) -> None:
|
||||
transformer = _tiny_transformer(fill=1.0)
|
||||
cb = EMACallback(decay=0.9, start_iter=10)
|
||||
cb.on_train_start(_Method(transformer), iteration=0)
|
||||
|
||||
# Drift transformer so that re-init has a visible effect.
|
||||
with torch.no_grad():
|
||||
transformer.weight.fill_(5.0)
|
||||
|
||||
cb.on_training_step_end(
|
||||
_Method(transformer), loss_dict={}, iteration=10
|
||||
)
|
||||
# First active step: shadow is re-initialized from the
|
||||
# current transformer (5.0) and *then* update() applies decay
|
||||
# against the same value, so shadow stays at 5.0.
|
||||
assert cb._ema_started
|
||||
assert torch.allclose(
|
||||
cb.student_ema.shadow["weight"],
|
||||
torch.full((2, 4), 5.0),
|
||||
)
|
||||
|
||||
def test_subsequent_step_applies_decay(self) -> None:
|
||||
transformer = _tiny_transformer(fill=2.0)
|
||||
cb = EMACallback(decay=0.9, start_iter=0)
|
||||
cb.on_train_start(_Method(transformer), iteration=0)
|
||||
|
||||
# Step 0: re-init at 2.0, then update against 2.0 → still 2.0.
|
||||
cb.on_training_step_end(
|
||||
_Method(transformer), loss_dict={}, iteration=0
|
||||
)
|
||||
# Step 1: drift transformer to 12.0, expect
|
||||
# shadow = 0.9 * 2.0 + 0.1 * 12.0 = 3.0.
|
||||
with torch.no_grad():
|
||||
transformer.weight.fill_(12.0)
|
||||
cb.on_training_step_end(
|
||||
_Method(transformer), loss_dict={}, iteration=1
|
||||
)
|
||||
assert torch.allclose(
|
||||
cb.student_ema.shadow["weight"],
|
||||
torch.full((2, 4), 3.0),
|
||||
atol=1e-6,
|
||||
)
|
||||
|
||||
def test_tracker_logs_decay(self) -> None:
|
||||
transformer = _tiny_transformer()
|
||||
tracker = _RecordingTracker()
|
||||
cb = EMACallback(decay=0.99, start_iter=0)
|
||||
method = _Method(transformer, tracker=tracker)
|
||||
cb.on_train_start(method, iteration=0)
|
||||
cb.on_training_step_end(method, loss_dict={}, iteration=0)
|
||||
|
||||
assert any(
|
||||
payload.get("ema/decay") == 0.99 and step == 0
|
||||
for payload, step in tracker.entries
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# C. ema_context
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEmaContext:
|
||||
|
||||
def test_passthrough_when_inactive(self) -> None:
|
||||
transformer = _tiny_transformer(fill=3.0)
|
||||
cb = EMACallback()
|
||||
# No on_train_start → student_ema is None.
|
||||
with cb.ema_context(transformer) as t:
|
||||
assert t is transformer
|
||||
assert torch.allclose(
|
||||
t.weight, torch.full((2, 4), 3.0)
|
||||
)
|
||||
|
||||
def test_swaps_weights_then_restores(self) -> None:
|
||||
transformer = _tiny_transformer(fill=1.0)
|
||||
cb = EMACallback(decay=0.0, start_iter=0)
|
||||
method = _Method(transformer)
|
||||
cb.on_train_start(method, iteration=0)
|
||||
|
||||
# decay=0 → after one step the shadow == current weights == 1.0.
|
||||
cb.on_training_step_end(method, loss_dict={}, iteration=0)
|
||||
# Drift transformer; ema_context should swap shadow (1.0) in
|
||||
# for the duration and restore the post-drift value (9.0).
|
||||
with torch.no_grad():
|
||||
transformer.weight.fill_(9.0)
|
||||
|
||||
with cb.ema_context(transformer) as t:
|
||||
assert torch.allclose(t.weight, torch.full((2, 4), 1.0))
|
||||
|
||||
assert torch.allclose(
|
||||
transformer.weight, torch.full((2, 4), 9.0)
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# D. State dict round-trip
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStateDict:
|
||||
|
||||
def test_state_dict_empty_before_train_start(self) -> None:
|
||||
cb = EMACallback()
|
||||
assert cb.state_dict() == {}
|
||||
|
||||
def test_round_trip_preserves_shadow_and_started_flag(self) -> None:
|
||||
transformer = _tiny_transformer(fill=4.0)
|
||||
cb = EMACallback(decay=0.5, start_iter=0)
|
||||
method = _Method(transformer)
|
||||
cb.on_train_start(method, iteration=0)
|
||||
cb.on_training_step_end(method, loss_dict={}, iteration=0)
|
||||
|
||||
state = cb.state_dict()
|
||||
assert "student_ema" in state
|
||||
assert state["ema_started"] is True
|
||||
|
||||
# Build a fresh callback and load.
|
||||
fresh = EMACallback(decay=0.5, start_iter=0)
|
||||
fresh.on_train_start(_Method(_tiny_transformer(fill=0.0)),
|
||||
iteration=0)
|
||||
# Sanity: fresh shadow != saved shadow before load.
|
||||
assert not torch.allclose(
|
||||
fresh.student_ema.shadow["weight"],
|
||||
cb.student_ema.shadow["weight"],
|
||||
)
|
||||
fresh.load_state_dict(state)
|
||||
assert fresh._ema_started is True
|
||||
assert torch.allclose(
|
||||
fresh.student_ema.shadow["weight"],
|
||||
cb.student_ema.shadow["weight"],
|
||||
)
|
||||
|
||||
def test_load_without_student_ema_only_sets_flag(self) -> None:
|
||||
cb = EMACallback()
|
||||
# student_ema is None — load must not attempt to assign shadow.
|
||||
cb.load_state_dict({"ema_started": True})
|
||||
assert cb._ema_started is True
|
||||
assert cb.student_ema is None
|
||||
@@ -1,170 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.grad_clip`.
|
||||
|
||||
Exercises ``GradNormClipCallback.on_before_optimizer_step`` against
|
||||
synthetic ``nn.Module`` targets with manually populated gradients.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.train.callbacks.grad_clip import GradNormClipCallback
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _RecordingTracker:
|
||||
"""Tracker stub that records every ``log`` call."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.entries: list[tuple[dict[str, Any], int]] = []
|
||||
|
||||
def log(self, payload: dict[str, Any], step: int) -> None:
|
||||
self.entries.append((payload, step))
|
||||
|
||||
|
||||
class _Method:
|
||||
"""Minimal stand-in for ``TrainingMethod``."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
targets: dict[str, torch.nn.Module],
|
||||
tracker: Any | None = None,
|
||||
) -> None:
|
||||
self._targets = targets
|
||||
self.tracker = tracker
|
||||
self.iter_seen: int | None = None
|
||||
|
||||
def get_grad_clip_targets(
|
||||
self, iteration: int
|
||||
) -> dict[str, torch.nn.Module]:
|
||||
self.iter_seen = iteration
|
||||
return self._targets
|
||||
|
||||
|
||||
def _make_module(*, grad_value: float, n: int = 4) -> torch.nn.Module:
|
||||
"""Return an ``nn.Linear`` whose grads are filled with ``grad_value``."""
|
||||
m = torch.nn.Linear(n, n, bias=False)
|
||||
m.weight.grad = torch.full_like(m.weight, fill_value=grad_value)
|
||||
return m
|
||||
|
||||
|
||||
def _grad_norm(module: torch.nn.Module) -> float:
|
||||
flat = torch.cat([p.grad.flatten() for p in module.parameters()])
|
||||
return float(flat.norm(2).item())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGradNormClipCallback:
|
||||
|
||||
def test_disabled_when_max_norm_non_positive(self) -> None:
|
||||
m = _make_module(grad_value=10.0)
|
||||
before = _grad_norm(m)
|
||||
|
||||
cb = GradNormClipCallback(max_grad_norm=0.0)
|
||||
method = _Method(targets={"m": m})
|
||||
cb.on_before_optimizer_step(method=method, iteration=0)
|
||||
|
||||
# No clipping applied; ``get_grad_clip_targets`` not consulted.
|
||||
assert _grad_norm(m) == before
|
||||
assert method.iter_seen is None
|
||||
|
||||
def test_large_grads_get_clipped(self) -> None:
|
||||
m = _make_module(grad_value=10.0)
|
||||
assert _grad_norm(m) > 1.0
|
||||
|
||||
cb = GradNormClipCallback(max_grad_norm=1.0)
|
||||
cb.on_before_optimizer_step(
|
||||
method=_Method(targets={"m": m}), iteration=0
|
||||
)
|
||||
|
||||
# After clipping the L2 norm should not exceed max_grad_norm
|
||||
# (allow a tiny epsilon for the +1e-6 in the clip helper).
|
||||
assert _grad_norm(m) <= 1.0 + 1e-4
|
||||
|
||||
def test_small_grads_unchanged(self) -> None:
|
||||
m = _make_module(grad_value=0.01)
|
||||
before = _grad_norm(m)
|
||||
assert before < 1.0
|
||||
|
||||
cb = GradNormClipCallback(max_grad_norm=1.0)
|
||||
cb.on_before_optimizer_step(
|
||||
method=_Method(targets={"m": m}), iteration=0
|
||||
)
|
||||
|
||||
# Clip coef >1 is clamped to 1, so values are preserved
|
||||
# (modulo the *1.0 multiply, which is exact for floats).
|
||||
assert abs(_grad_norm(m) - before) < 1e-6
|
||||
|
||||
def test_iteration_forwarded_to_targets(self) -> None:
|
||||
m = _make_module(grad_value=0.5)
|
||||
cb = GradNormClipCallback(max_grad_norm=1.0)
|
||||
method = _Method(targets={"m": m})
|
||||
cb.on_before_optimizer_step(method=method, iteration=42)
|
||||
assert method.iter_seen == 42
|
||||
|
||||
def test_tracker_logged_when_enabled(self) -> None:
|
||||
m = _make_module(grad_value=5.0)
|
||||
tracker = _RecordingTracker()
|
||||
cb = GradNormClipCallback(
|
||||
max_grad_norm=1.0, log_grad_norms=True
|
||||
)
|
||||
cb.on_before_optimizer_step(
|
||||
method=_Method(targets={"layer": m}, tracker=tracker),
|
||||
iteration=7,
|
||||
)
|
||||
assert len(tracker.entries) == 1
|
||||
payload, step = tracker.entries[0]
|
||||
assert step == 7
|
||||
assert "grad_norm/layer" in payload
|
||||
assert payload["grad_norm/layer"] > 0.0
|
||||
|
||||
def test_tracker_not_logged_when_disabled(self) -> None:
|
||||
m = _make_module(grad_value=5.0)
|
||||
tracker = _RecordingTracker()
|
||||
cb = GradNormClipCallback(
|
||||
max_grad_norm=1.0, log_grad_norms=False
|
||||
)
|
||||
cb.on_before_optimizer_step(
|
||||
method=_Method(targets={"m": m}, tracker=tracker),
|
||||
iteration=0,
|
||||
)
|
||||
assert tracker.entries == []
|
||||
|
||||
def test_no_tracker_does_not_raise(self) -> None:
|
||||
m = _make_module(grad_value=5.0)
|
||||
cb = GradNormClipCallback(max_grad_norm=1.0, log_grad_norms=True)
|
||||
# Method without a tracker attribute at all.
|
||||
|
||||
class _BareMethod:
|
||||
|
||||
def get_grad_clip_targets(
|
||||
self, iteration: int
|
||||
) -> dict[str, torch.nn.Module]:
|
||||
return {"m": m}
|
||||
|
||||
cb.on_before_optimizer_step(method=_BareMethod(), iteration=0)
|
||||
# No assertion — must simply not raise.
|
||||
|
||||
def test_multiple_targets_each_logged(self) -> None:
|
||||
targets = {
|
||||
"head": _make_module(grad_value=4.0),
|
||||
"tail": _make_module(grad_value=8.0),
|
||||
}
|
||||
tracker = _RecordingTracker()
|
||||
cb = GradNormClipCallback(max_grad_norm=1.0)
|
||||
cb.on_before_optimizer_step(
|
||||
method=_Method(targets=targets, tracker=tracker),
|
||||
iteration=1,
|
||||
)
|
||||
keys = {next(iter(p)) for p, _ in tracker.entries}
|
||||
assert keys == {"grad_norm/head", "grad_norm/tail"}
|
||||
@@ -1,234 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only unit tests for :mod:`fastvideo.train.callbacks.validation`.
|
||||
|
||||
Covers the parts of ``ValidationCallback`` that don't need a real
|
||||
pipeline or distributed init:
|
||||
|
||||
* constructor type coercions and defaults,
|
||||
* ``on_validation_begin`` gating logic (every_steps + modulo),
|
||||
* ``_find_ema_callback`` lookup via ``_callback_dict``,
|
||||
* ``state_dict`` / ``load_state_dict`` rng round-trip.
|
||||
|
||||
The heavy ``_run_validation`` path needs a real diffusion pipeline plus
|
||||
distributed init and is exercised by Phase 2/3 tests.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.train.callbacks.callback import CallbackDict
|
||||
from fastvideo.train.callbacks.ema import EMACallback
|
||||
from fastvideo.train.callbacks.validation import ValidationCallback
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
_PIPE_TARGET = "fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline"
|
||||
|
||||
|
||||
def _make_callback(
|
||||
*,
|
||||
every_steps: int = 100,
|
||||
sampling_steps: list[int] | None = None,
|
||||
guidance_scale: float | None = None,
|
||||
num_frames: int | None = None,
|
||||
sampling_timesteps: list[int] | None = None,
|
||||
output_dir: str | None = None,
|
||||
) -> ValidationCallback:
|
||||
return ValidationCallback(
|
||||
pipeline_target=_PIPE_TARGET,
|
||||
dataset_file="/tmp/does_not_exist.json",
|
||||
every_steps=every_steps,
|
||||
sampling_steps=sampling_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
num_frames=num_frames,
|
||||
sampling_timesteps=sampling_timesteps,
|
||||
output_dir=output_dir,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A. Constructor coercions / defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConstructor:
|
||||
|
||||
def test_defaults(self) -> None:
|
||||
cb = _make_callback()
|
||||
assert cb.pipeline_target == _PIPE_TARGET
|
||||
assert cb.dataset_file == "/tmp/does_not_exist.json"
|
||||
assert cb.every_steps == 100
|
||||
assert cb.sampling_steps == [40]
|
||||
assert cb.guidance_scale is None
|
||||
assert cb.num_frames is None
|
||||
assert cb.sampling_timesteps is None
|
||||
assert cb.output_dir is None
|
||||
# Lazy fields not yet populated.
|
||||
assert cb._pipeline is None
|
||||
assert cb._sampling_param is None
|
||||
assert cb.validation_random_generator is None
|
||||
|
||||
def test_string_inputs_are_coerced(self) -> None:
|
||||
# YAML often produces strings for numeric fields; the
|
||||
# constructor must coerce them.
|
||||
cb = ValidationCallback(
|
||||
pipeline_target=_PIPE_TARGET,
|
||||
dataset_file="x.json",
|
||||
every_steps="50", # type: ignore[arg-type]
|
||||
sampling_steps=["20", "40"], # type: ignore[arg-type]
|
||||
guidance_scale="4.5", # type: ignore[arg-type]
|
||||
num_frames="77", # type: ignore[arg-type]
|
||||
sampling_timesteps=["1000", "500"],
|
||||
)
|
||||
assert cb.every_steps == 50
|
||||
assert cb.sampling_steps == [20, 40]
|
||||
assert cb.guidance_scale == 4.5
|
||||
assert cb.num_frames == 77
|
||||
assert cb.sampling_timesteps == [1000, 500]
|
||||
|
||||
def test_pipeline_kwargs_collected(self) -> None:
|
||||
cb = ValidationCallback(
|
||||
pipeline_target=_PIPE_TARGET,
|
||||
dataset_file="x.json",
|
||||
extra_arg=123,
|
||||
another="value",
|
||||
)
|
||||
# Unknown kwargs are stashed for the pipeline factory.
|
||||
assert cb.pipeline_kwargs == {
|
||||
"extra_arg": 123,
|
||||
"another": "value",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# B. on_validation_begin gating
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _NoRunValidation(ValidationCallback):
|
||||
"""Subclass that records ``_run_validation`` calls instead of
|
||||
actually running them — lets us assert the gating logic without a
|
||||
real pipeline."""
|
||||
|
||||
def __init__(self, **kwargs) -> None:
|
||||
super().__init__(**kwargs)
|
||||
self.run_calls: list[int] = []
|
||||
|
||||
def _run_validation(self, method, step: int) -> None: # type: ignore[override]
|
||||
self.run_calls.append(step)
|
||||
|
||||
|
||||
def _make_recording(**kwargs) -> _NoRunValidation:
|
||||
return _NoRunValidation(
|
||||
pipeline_target=_PIPE_TARGET,
|
||||
dataset_file="x.json",
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class TestOnValidationBegin:
|
||||
|
||||
def test_skipped_when_every_steps_zero(self) -> None:
|
||||
cb = _make_recording(every_steps=0)
|
||||
cb.on_validation_begin(method=None, iteration=0)
|
||||
cb.on_validation_begin(method=None, iteration=1000)
|
||||
assert cb.run_calls == []
|
||||
|
||||
def test_skipped_on_off_iter(self) -> None:
|
||||
cb = _make_recording(every_steps=50)
|
||||
cb.on_validation_begin(method=None, iteration=49)
|
||||
cb.on_validation_begin(method=None, iteration=51)
|
||||
assert cb.run_calls == []
|
||||
|
||||
def test_runs_on_match(self) -> None:
|
||||
cb = _make_recording(every_steps=50)
|
||||
cb.on_validation_begin(method=None, iteration=50)
|
||||
cb.on_validation_begin(method=None, iteration=100)
|
||||
assert cb.run_calls == [50, 100]
|
||||
|
||||
def test_iter_zero_runs(self) -> None:
|
||||
# 0 % anything == 0 → step 0 fires (matches existing
|
||||
# validation behavior used by ValidationCallback consumers).
|
||||
cb = _make_recording(every_steps=50)
|
||||
cb.on_validation_begin(method=None, iteration=0)
|
||||
assert cb.run_calls == [0]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# C. _find_ema_callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFindEmaCallback:
|
||||
|
||||
def test_returns_none_without_callback_dict(self) -> None:
|
||||
cb = _make_callback()
|
||||
# _callback_dict is not set on bare instances.
|
||||
assert cb._find_ema_callback() is None
|
||||
|
||||
def test_returns_none_when_no_ema_registered(self) -> None:
|
||||
cb = _make_callback()
|
||||
cb_dict = CallbackDict({}, training_config=object())
|
||||
cb._callback_dict = cb_dict
|
||||
assert cb._find_ema_callback() is None
|
||||
|
||||
def test_finds_ema_callback(self) -> None:
|
||||
cb = _make_callback()
|
||||
cb_dict = CallbackDict({}, training_config=object())
|
||||
ema = EMACallback(decay=0.99)
|
||||
cb_dict._callbacks["ema"] = ema
|
||||
cb_dict._callbacks["validation"] = cb
|
||||
cb._callback_dict = cb_dict
|
||||
|
||||
found = cb._find_ema_callback()
|
||||
assert found is ema
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# D. state_dict / load_state_dict (rng round-trip)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStateDict:
|
||||
|
||||
def test_state_dict_empty_without_generator(self) -> None:
|
||||
cb = _make_callback()
|
||||
# validation_random_generator is None until on_train_start.
|
||||
assert cb.state_dict() == {}
|
||||
|
||||
def test_round_trip_preserves_rng_state(self) -> None:
|
||||
cb = _make_callback()
|
||||
gen = torch.Generator(device="cpu").manual_seed(123)
|
||||
# Advance RNG so a default-init generator on the receiving
|
||||
# side is observably different.
|
||||
for _ in range(5):
|
||||
torch.randn(4, generator=gen)
|
||||
cb.validation_random_generator = gen
|
||||
|
||||
state = cb.state_dict()
|
||||
assert "validation_rng" in state
|
||||
|
||||
# Receiver: fresh generator with a different seed.
|
||||
fresh = _make_callback()
|
||||
fresh.validation_random_generator = (
|
||||
torch.Generator(device="cpu").manual_seed(999)
|
||||
)
|
||||
fresh.load_state_dict(state)
|
||||
|
||||
# After load, both generators draw the same next sample.
|
||||
a = torch.randn(8, generator=cb.validation_random_generator)
|
||||
b = torch.randn(8, generator=fresh.validation_random_generator)
|
||||
assert torch.equal(a, b)
|
||||
|
||||
def test_load_without_generator_is_noop(self) -> None:
|
||||
cb = _make_callback()
|
||||
# Generator is None: load must not raise even when state has
|
||||
# an rng entry.
|
||||
cb.load_state_dict(
|
||||
{"validation_rng": torch.tensor([1, 2, 3], dtype=torch.uint8)}
|
||||
)
|
||||
assert cb.validation_random_generator is None
|
||||
@@ -1,360 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only unit tests for :mod:`fastvideo.train.utils.checkpoint`.
|
||||
|
||||
Covers the pure-Python portions of the checkpoint manager: name
|
||||
parsing, resume-path resolution, metadata round-trip, rolling-delete
|
||||
cleanup, the ``_is_stateful`` predicate, and the ``maybe_save`` gating
|
||||
logic. Code paths that touch DCP (``dcp.save`` / ``dcp.load``) and
|
||||
CUDA RNG snapshots are intentionally not covered here — those need a
|
||||
GPU runner and will be tested in later phases.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.train.utils.checkpoint import (
|
||||
CheckpointConfig,
|
||||
CheckpointManager,
|
||||
_find_latest_checkpoint,
|
||||
_is_stateful,
|
||||
_parse_step_from_dir,
|
||||
_resolve_resume_checkpoint,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_checkpoint_dir(
|
||||
output_dir: Path,
|
||||
step: int,
|
||||
*,
|
||||
with_dcp: bool = True,
|
||||
) -> Path:
|
||||
"""Create a fake ``checkpoint-<step>/dcp`` directory tree."""
|
||||
ckpt_dir = output_dir / f"checkpoint-{step}"
|
||||
ckpt_dir.mkdir(parents=True, exist_ok=True)
|
||||
if with_dcp:
|
||||
(ckpt_dir / "dcp").mkdir(exist_ok=True)
|
||||
return ckpt_dir
|
||||
|
||||
|
||||
def _make_manager(
|
||||
tmp_path: Path,
|
||||
*,
|
||||
save_steps: int = 0,
|
||||
keep_last: int = 0,
|
||||
raw_config: dict[str, Any] | None = None,
|
||||
) -> CheckpointManager:
|
||||
"""Build a minimal ``CheckpointManager`` for tests that don't touch DCP."""
|
||||
return CheckpointManager(
|
||||
method=None,
|
||||
dataloader=None,
|
||||
output_dir=str(tmp_path),
|
||||
config=CheckpointConfig(save_steps=save_steps, keep_last=keep_last),
|
||||
raw_config=raw_config,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A. _is_stateful predicate
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _Full:
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def load_state_dict(self, sd: dict[str, Any]) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class _MissingStateDict:
|
||||
|
||||
def load_state_dict(self, sd: dict[str, Any]) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class _MissingLoad:
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return {}
|
||||
|
||||
|
||||
def test_is_stateful_true_for_full_object() -> None:
|
||||
assert _is_stateful(_Full()) is True
|
||||
|
||||
|
||||
def test_is_stateful_false_when_missing_state_dict() -> None:
|
||||
assert _is_stateful(_MissingStateDict()) is False
|
||||
|
||||
|
||||
def test_is_stateful_false_when_missing_load_state_dict() -> None:
|
||||
assert _is_stateful(_MissingLoad()) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# B. _parse_step_from_dir
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_parse_step_valid(tmp_path: Path) -> None:
|
||||
assert _parse_step_from_dir(tmp_path / "checkpoint-100") == 100
|
||||
|
||||
|
||||
def test_parse_step_zero(tmp_path: Path) -> None:
|
||||
assert _parse_step_from_dir(tmp_path / "checkpoint-0") == 0
|
||||
|
||||
|
||||
def test_parse_step_invalid_raises(tmp_path: Path) -> None:
|
||||
with pytest.raises(ValueError, match="Invalid checkpoint directory"):
|
||||
_parse_step_from_dir(tmp_path / "not-a-checkpoint")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# C. _find_latest_checkpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_find_latest_returns_none_on_nonexistent_dir(tmp_path: Path) -> None:
|
||||
assert _find_latest_checkpoint(tmp_path / "missing") is None
|
||||
|
||||
|
||||
def test_find_latest_returns_none_on_empty_dir(tmp_path: Path) -> None:
|
||||
assert _find_latest_checkpoint(tmp_path) is None
|
||||
|
||||
|
||||
def test_find_latest_returns_largest_step(tmp_path: Path) -> None:
|
||||
_make_checkpoint_dir(tmp_path, 10)
|
||||
_make_checkpoint_dir(tmp_path, 200)
|
||||
_make_checkpoint_dir(tmp_path, 50)
|
||||
latest = _find_latest_checkpoint(tmp_path)
|
||||
assert latest is not None
|
||||
assert latest.name == "checkpoint-200"
|
||||
|
||||
|
||||
def test_find_latest_skips_dirs_without_dcp_subdir(tmp_path: Path) -> None:
|
||||
# checkpoint-10 is "corrupted" — has no dcp/ subdir, must be skipped.
|
||||
_make_checkpoint_dir(tmp_path, 10, with_dcp=False)
|
||||
_make_checkpoint_dir(tmp_path, 5, with_dcp=True)
|
||||
latest = _find_latest_checkpoint(tmp_path)
|
||||
assert latest is not None
|
||||
assert latest.name == "checkpoint-5"
|
||||
|
||||
|
||||
def test_find_latest_skips_non_checkpoint_dirs(tmp_path: Path) -> None:
|
||||
(tmp_path / "logs").mkdir()
|
||||
(tmp_path / "wandb").mkdir()
|
||||
(tmp_path / "some_file.txt").write_text("noise")
|
||||
_make_checkpoint_dir(tmp_path, 7)
|
||||
latest = _find_latest_checkpoint(tmp_path)
|
||||
assert latest is not None
|
||||
assert latest.name == "checkpoint-7"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# D. _resolve_resume_checkpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_latest_with_no_checkpoints_returns_none(
|
||||
tmp_path: Path) -> None:
|
||||
out = tmp_path / "outputs"
|
||||
out.mkdir()
|
||||
assert _resolve_resume_checkpoint("latest", output_dir=str(out)) is None
|
||||
|
||||
|
||||
def test_resolve_latest_returns_latest_checkpoint(tmp_path: Path) -> None:
|
||||
_make_checkpoint_dir(tmp_path, 30)
|
||||
_make_checkpoint_dir(tmp_path, 10)
|
||||
resolved = _resolve_resume_checkpoint("latest", output_dir=str(tmp_path))
|
||||
assert resolved is not None
|
||||
assert resolved.name == "checkpoint-30"
|
||||
|
||||
|
||||
def test_resolve_explicit_checkpoint_dir(tmp_path: Path) -> None:
|
||||
ckpt = _make_checkpoint_dir(tmp_path, 42)
|
||||
resolved = _resolve_resume_checkpoint(str(ckpt),
|
||||
output_dir=str(tmp_path))
|
||||
assert resolved is not None
|
||||
assert resolved.name == "checkpoint-42"
|
||||
|
||||
|
||||
def test_resolve_dcp_subdir_returns_parent_checkpoint(tmp_path: Path) -> None:
|
||||
ckpt = _make_checkpoint_dir(tmp_path, 42)
|
||||
dcp_path = ckpt / "dcp"
|
||||
resolved = _resolve_resume_checkpoint(str(dcp_path),
|
||||
output_dir=str(tmp_path))
|
||||
assert resolved is not None
|
||||
assert resolved.name == "checkpoint-42"
|
||||
|
||||
|
||||
def test_resolve_output_dir_returns_latest(tmp_path: Path) -> None:
|
||||
out = tmp_path / "outputs"
|
||||
out.mkdir()
|
||||
_make_checkpoint_dir(out, 100)
|
||||
_make_checkpoint_dir(out, 50)
|
||||
resolved = _resolve_resume_checkpoint(str(out), output_dir=str(tmp_path))
|
||||
assert resolved is not None
|
||||
assert resolved.name == "checkpoint-100"
|
||||
|
||||
|
||||
def test_resolve_nonexistent_path_raises(tmp_path: Path) -> None:
|
||||
with pytest.raises(FileNotFoundError):
|
||||
_resolve_resume_checkpoint(str(tmp_path / "missing"),
|
||||
output_dir=str(tmp_path))
|
||||
|
||||
|
||||
def test_resolve_checkpoint_without_dcp_raises(tmp_path: Path) -> None:
|
||||
ckpt = _make_checkpoint_dir(tmp_path, 5, with_dcp=False)
|
||||
with pytest.raises(FileNotFoundError, match="dcp"):
|
||||
_resolve_resume_checkpoint(str(ckpt), output_dir=str(tmp_path))
|
||||
|
||||
|
||||
def test_resolve_unknown_dir_raises(tmp_path: Path) -> None:
|
||||
"""A dir that is neither a checkpoint nor an output_dir-with-checkpoints."""
|
||||
bogus = tmp_path / "bogus"
|
||||
bogus.mkdir()
|
||||
with pytest.raises(ValueError, match="Could not resolve"):
|
||||
_resolve_resume_checkpoint(str(bogus), output_dir=str(tmp_path))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# E. metadata read/write
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_write_metadata_roundtrip_with_step(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path)
|
||||
ckpt_dir = _make_checkpoint_dir(tmp_path, 7)
|
||||
mgr._write_metadata(ckpt_dir, step=7)
|
||||
loaded = CheckpointManager.load_metadata(ckpt_dir)
|
||||
assert loaded == {"step": 7}
|
||||
|
||||
|
||||
def test_write_metadata_includes_raw_config(tmp_path: Path) -> None:
|
||||
raw = {
|
||||
"models": {
|
||||
"student": {
|
||||
"_target_": "X"
|
||||
}
|
||||
},
|
||||
"training": {
|
||||
"distributed": {
|
||||
"num_gpus": 4
|
||||
}
|
||||
},
|
||||
}
|
||||
mgr = _make_manager(tmp_path, raw_config=raw)
|
||||
ckpt_dir = _make_checkpoint_dir(tmp_path, 7)
|
||||
mgr._write_metadata(ckpt_dir, step=7)
|
||||
loaded = CheckpointManager.load_metadata(ckpt_dir)
|
||||
assert loaded["step"] == 7
|
||||
assert loaded["config"] == raw
|
||||
|
||||
|
||||
def test_load_metadata_raises_on_missing_file(tmp_path: Path) -> None:
|
||||
ckpt_dir = _make_checkpoint_dir(tmp_path, 7)
|
||||
# No metadata.json written.
|
||||
with pytest.raises(FileNotFoundError, match="metadata"):
|
||||
CheckpointManager.load_metadata(ckpt_dir)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# F. _cleanup_old_checkpoints (rolling delete)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_cleanup_keep_last_zero_is_noop(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, keep_last=0)
|
||||
for step in (1, 2, 3):
|
||||
_make_checkpoint_dir(tmp_path, step)
|
||||
mgr._cleanup_old_checkpoints()
|
||||
remaining = sorted(p.name for p in tmp_path.iterdir())
|
||||
assert remaining == ["checkpoint-1", "checkpoint-2", "checkpoint-3"]
|
||||
|
||||
|
||||
def test_cleanup_keeps_newest_when_over_limit(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, keep_last=2)
|
||||
for step in (1, 5, 10, 50, 100):
|
||||
_make_checkpoint_dir(tmp_path, step)
|
||||
mgr._cleanup_old_checkpoints()
|
||||
remaining = sorted(p.name for p in tmp_path.iterdir())
|
||||
assert remaining == ["checkpoint-100", "checkpoint-50"]
|
||||
|
||||
|
||||
def test_cleanup_no_op_when_under_limit(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, keep_last=3)
|
||||
for step in (1, 2):
|
||||
_make_checkpoint_dir(tmp_path, step)
|
||||
mgr._cleanup_old_checkpoints()
|
||||
remaining = sorted(p.name for p in tmp_path.iterdir())
|
||||
assert remaining == ["checkpoint-1", "checkpoint-2"]
|
||||
|
||||
|
||||
def test_cleanup_skips_non_checkpoint_dirs(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, keep_last=1)
|
||||
for step in (1, 2, 3):
|
||||
_make_checkpoint_dir(tmp_path, step)
|
||||
(tmp_path / "logs").mkdir()
|
||||
(tmp_path / "wandb").mkdir()
|
||||
mgr._cleanup_old_checkpoints()
|
||||
remaining = sorted(p.name for p in tmp_path.iterdir())
|
||||
assert remaining == ["checkpoint-3", "logs", "wandb"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# G. maybe_save gating logic
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _record_save_calls(mgr: CheckpointManager) -> list[int]:
|
||||
"""Replace ``mgr.save`` with a recorder that mimics the side effect
|
||||
of advancing ``_last_saved_step`` so dedup logic still works."""
|
||||
calls: list[int] = []
|
||||
|
||||
def fake_save(step: int) -> None:
|
||||
calls.append(step)
|
||||
mgr._last_saved_step = step
|
||||
|
||||
mgr.save = fake_save # type: ignore[method-assign]
|
||||
return calls
|
||||
|
||||
|
||||
def test_maybe_save_skipped_when_save_steps_is_zero(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, save_steps=0)
|
||||
calls = _record_save_calls(mgr)
|
||||
mgr.maybe_save(step=10)
|
||||
mgr.maybe_save(step=100)
|
||||
assert calls == []
|
||||
|
||||
|
||||
def test_maybe_save_skipped_when_step_not_on_interval(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, save_steps=10)
|
||||
calls = _record_save_calls(mgr)
|
||||
for step in (1, 5, 9, 11, 15):
|
||||
mgr.maybe_save(step=step)
|
||||
assert calls == []
|
||||
|
||||
|
||||
def test_maybe_save_dedupes_on_repeated_call(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, save_steps=10)
|
||||
calls = _record_save_calls(mgr)
|
||||
mgr.maybe_save(step=20)
|
||||
mgr.maybe_save(step=20)
|
||||
mgr.maybe_save(step=20)
|
||||
assert calls == [20]
|
||||
|
||||
|
||||
def test_maybe_save_triggers_on_each_interval(tmp_path: Path) -> None:
|
||||
mgr = _make_manager(tmp_path, save_steps=10)
|
||||
calls = _record_save_calls(mgr)
|
||||
for step in range(1, 41):
|
||||
mgr.maybe_save(step=step)
|
||||
assert calls == [10, 20, 30, 40]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user