Compare commits

...
Author SHA1 Message Date
Will LinandClaude Fable 5 0a861031c7 [docs] H3 parallel VAE: document the compiled-decoder cross-process determinism caveat
With enable_torch_compile_vae (#1734, opt-in) inductor autotunes kernels per
process, so chunk decodes on other ranks differ from the serial rank's decode
the way two serial processes differ (GB200 @124f: max 63/255 on <0.5% of
pixels, mean ~1e-2/255; audio and chunk 0 bit-identical). Eager decoder (the
default) stays bitwise-equal to serial decode_to_pixels - measured, both
strategies, x3, 124f+345f.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 18:17:28 +00:00
Will LinandClaude Fable 5 f08c5ee8af [perf] H3 parallel VAE: overlap first decodes with the meta rendezvous; assembly on a side stream
Two schedule fixes sized from the first GB200 tray measurements (job 2659,
serial 7.7s/21.2s at 124f/345f):

1. Every rank now decodes its round-0 chunk BEFORE the metadata broadcast.
   Non-leader ranks previously blocked on the broadcast until the leader
   finished chunk 0, serializing a full extra chunk-decode into round 0
   (visible as 1.9x instead of ~2.6x at 7 chunks / 4 ranks). Same reorder
   on the encode path.

2. The leader's per-chunk joining work (blend, denormalize, clamp, output
   copies) moves to a dedicated CUDA side stream. It depends only on
   already-gathered segments, but on the main stream it delayed the
   leader's next-round decode and therefore every rank's next collective
   (~0.1s/chunk on the critical path). Gathered storage is pinned to the
   assembly stream via record_stream; the driver drains the stream in a
   finally so an exception cannot leave an in-flight DMA into the output
   buffer. Stream placement does not change op order or values, so the
   bitwise-parity contract is untouched.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:59:18 +00:00
Will LinandClaude Fable 5 755f4a7967 [docs] schema parity inventory: classify vae_parallel_* (+ the stack's unclassified VSA_tile_size)
vae_parallel_decode / vae_parallel_encode / vae_parallel_decode_strategy are
model-specific optimization knobs (compatibility_only, like VSA_sparsity).
VSA_tile_size came in with the merged tile-64 route without an inventory
entry and failed test_fastvideo_args_fields_are_classified on the whole
stack; classify it the same way. The remaining pipeline_config inventory
gaps (image_encoder_precisions, ...) predate this branch and are left for
the owning PRs.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:53:41 +00:00
Will LinandClaude Fable 5 d543a67b10 [perf] MiniMax-H3 VAE: SP-rank-parallel chunk decode + reference-clip encode (opt-in)
Under SP>1 the H3 video VAE decoded all temporal chunks serially on the
output rank while the other ranks idled (#1703's gate), and every rank
encoded the full reference video redundantly. Chunk decodes and clip
encodes have no cross-chunk data dependency - only the joining (overlap
blend, trim, denormalize, moment concat) is sequential - so both are
round-robined across the sequence-parallel ranks:

- fastvideo/models/vaes/minimax_h3_parallel.py: decode_to_pixels_parallel
  gathers each round's decoded segments (body+halo tail, one contiguous
  slice per chunk) to the SP group's first rank via NCCL gather (or
  all_gather, FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY), which replays the
  serial blend/trim/denormalize/copy semantics with the same VAE methods -
  bitwise-equal to serial decode_to_pixels by construction. Placeholder
  rounds keep collective participation uniform; the leader decodes chunk 0
  first and broadcasts dtype/shape metadata so placeholders never guess the
  autocast dtype. encode_pixels_parallel all-gathers per-clip moments
  (latent-sized) so every rank keeps the identical full posterior,
  preserving the all-ranks-hold-latents contract.
- decoding stage: output gate moves from world rank 0 to the SP group's
  first rank (identical in the single-group e2e case; correct for the
  trainer validation callback, which consumes each group leader's batch);
  with vae_parallel_decode every rank enters the decode body so no
  rank-dependent branch guards the collectives.
- latent preparation: opt-in clip-parallel reference encode on the same
  seam (vae_parallel_encode).
- knobs: FastVideoArgs.vae_parallel_decode/encode (+ --vae-parallel-decode,
  --vae-parallel-encode, FASTVIDEO_VAE_PARALLEL_DECODE/ENCODE env
  parse-once adapters), default OFF.
- _copy_chunk_pixels factored out of _decode_to_pixels so serial and
  parallel share one output-copy path (behavior unchanged).

Tests: threaded fake-group CPU suite drives the real SPMD functions
end-to-end (world sizes 2-5, both strategies, pad/blend/trim geometries,
token_drop=0, batched slicing, placeholder rounds) bit-exact vs the serial
APIs; GPU regression (torchrun world>1 gated) asserts bitwise parity under
fp16 autocast with real NCCL plus repeat-determinism x3.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:38:38 +00:00
Will LinandClaude Fable 5 741aa8d289 [bugfix] logger: info_once crashed on the patched process-aware info (duplicate stacklevel)
_print_info_once passes stacklevel=2 into logger.info, and init_logger's
patched _info passed its own stacklevel=2 positionally into logger.log on
top of the caller's kwarg -> TypeError on every info_once call. Honor an
explicit stacklevel instead of passing the keyword twice.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:38:18 +00:00
Will LinandClaude Fable 5 ad9cd63122 [bugfix] H3 conditioner: no_grad instead of inference_mode - FSDP2-sharded encode crashed (#1732 default-path gap)
Carried blocker fix discovered while benchmarking this stack (GB200 job 2594,
all four legs): with text_encoder_cpu_offload=True - the FastVideoArgs DEFAULT
and the shipped H3 example configuration - TextEncoderLoader FSDP2-fully_shards
the H3 conditioner, and #1732's @torch.inference_mode() on
MiniMaxH3Qwen3VLConditioner.encode_ids then kills the very first encode:

  File torch/distributed/fsdp/_fully_shard/_fsdp_param_group.py, in
  wait_for_unshard: with torch.autograd._unsafe_preserve_version_counter(t):
  RuntimeError: Inference tensors do not track version counter.

FSDP2's lazy unshard runs inside the inference-mode region, so its all-gather
tensors are inference tensors, and the version-counter preservation hook
cannot read t._version. The PR's own benchmarks ran with offload disabled,
which is exactly the unexercised-default-path gap called out as finding 2 of
the pr1732 review (there for FP8; the same gap bites plain bf16 via
inference_mode).

Fix: @torch.no_grad() instead. It frees the same activation memory (the
-262 MiB claim comes from the early-return slim contract, not from
inference_mode's bookkeeping), is fully FSDP2-compatible, and as a bonus
retires review finding 7: prompt_embeds are ordinary tensors again, so any
future on-the-fly-conditioning training can backprop through them without a
clone at the stage boundary.

Verified: 1x GPU FastH3 leg boots and generates after this change (leg reruns
on GB200); the 1732 unit suites (truncation + checkpoint-fp8) still pass.
Note: fastvideo/tests/stages/test_text_encoding.py has 7 pre-existing failures
that reproduce byte-identically on plain origin/main 56d4a6074 (unrelated
generic-stage tests; not introduced by the stack, verified by A/B).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 10:49:47 +00:00
Will LinandClaude Fable 5 68e6ffca9e [bugfix] H3 VAE: clone the reduce-overhead stitched canvas at the tile-driver returns (#1734 review F1)
Carried blocker fix for PR #1734 on this integration stack, per pr1734.md
finding F1 (reproduced on GB200/torch 2.12): with
@torch.compile(mode="reduce-overhead") on _stitch_tiles, the stitched canvas
is a CUDA-graph static buffer, and the collect-then-torch.cat consumers -
_decode (via _decode_chunks), _encode, _encode_pixels, encode_keyframe - hold
each chunk/clip result across the next _stitch_tiles replay, which overwrites
the pooled storage. First tiled decode() with >=2 temporal chunks (any real
video) and tiled encode() of >17 frames raised:

  RuntimeError: Error: accessing tensor output of CUDAGraphs that has been
  overwritten by a subsequent run. ... line ..., in _stitch_tiles:
  return torch.cat(result_rows, dim=-2)

Fix: .clone() the stitch output at the two eager call sites (_encode_clip
tiled return, _decode_clip tiled return) so no cudagraph-owned storage
escapes the tile driver; the clone is read before the next replay, so it is
race-free by construction. The streaming _decode_to_pixels path was already
safe (full copy-out per chunk before the next decode) and stays correct at
one extra D2D copy per chunk (~tens of microseconds vs the decode compute).
This follows the option (a) recommendation in the review; reduce-overhead is
retained on _stitch_tiles/_project_decoder_tile.

Tests:
- test_decode_clip_emits_tiled_stage_ranges updated: the tile driver now
  returns a caller-owned copy (is-not + equal), pinning the ownership
  contract at the mock level.
- New CUDA regression gate (the review's must-add test):
  test_tiled_decode_and_encode_survive_cudagraph_buffer_reuse_on_cuda - real
  tiled decode() (2 temporal chunks, 2x2 spatial tiles) and encode()/
  encode_pixels() (2 clips) with _stitch_tiles unmocked, asserting bitwise
  repeat-consistency plus value parity against the fully eager tile helpers
  (via _torchdynamo_orig_callable). Red on the unfixed merge (cudagraph
  overwrite RuntimeError on the first decode), green with this fix.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 10:33:06 +00:00
Will Lin 64cdcf6be4 Merge PR #1731 head (9713ea127) on top of the stack (the target of this branch)
[feat] FastH3 few-step preview: VSA-H3 64-token tile option on the native
Triton block-sparse path, run-level --VSA-tile-size plumbing, opt-in sm_100a
CUDA forward route (FASTVIDEO_VSA_SM100A=1), basic_fasth3.py example with the
corrected 5-point-grid = 4-forward schedule (num_inference_steps=5), and the
FastVideo-Minimax-FastH3-Preview-v0.1 release name.

Clean textual merge; the predicted 1734<->1731 conflict in
minimax_h3_denoising.py resolved automatically and was verified by hand:
1731's vsa_tile_size plumbing sits in the metadata-builder preamble
(L149-152, L177) while 1734's edits are the import line, the stage docstring,
and the profiler_region+nvtx_range context line (L157) - the merged loop keeps
both, plus main's cudagraph_mark_step_begin contract (L184).
video_sparse_attn_h3.py and the metadata tests are 1731-only files on this
stack (no sibling edits).
2026-08-21 10:27:49 +00:00
Will Lin aa0d98a6b8 Merge PR #1734 head (dca423fd3) into integration/h3-perf-stack
[perf] MiniMax-H3 VAE decode optimization: _compile_conditions for the video +
audio VAE decoders (makes enable_torch_compile_vae effective for H3),
torch.compile on _stitch_tiles/_project_decoder_tile, VAE attention through the
FastVideo selector (TORCH_SDPA/FLASH_ATTN), opt-in FASTVIDEO_NVTX_PROFILE
ranges, per-worker attention-backend receipt.

Clean textual merge; semantic overlaps verified by hand:
- minimax_h3_conditioning.py: 1732's rewritten _encode_fl2va/_encode_ref2va
  kept their names, 1734's nvtx_range wrap composes over them (checked).
- minimax_h3.py: 1735 fusion routing + 1734 per-block nvtx_range coexist
  (fusions at attention/mlp/modulate seams, nvtx at the block loop).
- component_loader.py: 1732 quant plumbing at load_model (~L369) vs 1734
  backend receipt (~L1103) - disjoint.
- envs.py: FASTVIDEO_NVTX_PROFILE (1734) + FASTVIDEO_MINIMAX_H3_FUSIONS (1735)
  are distinct additions.
- minimax_h3_video.py: 1734 was authored on a base that already contains the
  merged #1703 (incl. pr1703-fixes content), so the 1734-vs-1703 conflict the
  reviews predicted was pre-resolved by the author's rebase.

KNOWN CARRIED DEFECT at this point in the stack: review finding F1 of
pr1734.md - reduce-overhead CUDA-graph outputs of _stitch_tiles escape into
collect-then-cat consumers (_decode/_encode/_encode_pixels), crashing tiled
eager decode/encode on CUDA. Fixed in a follow-up commit on this branch.
2026-08-21 10:26:23 +00:00
Will Lin 93b03bc14d Merge PR #1735 head (cbab605ef) into integration/h3-perf-stack
[perf] Opt-in MiniMax-H3 Sol-Engine Triton fusions (FASTVIDEO_MINIMAX_H3_FUSIONS):
fused rmsnorm+modulate, residual+gate+rmsnorm+modulate, per-head qknorm+partial
RoPE, and packed SwiGLU. Default-off; fused path is disclosed NON-PARITY
numerics (same-seed decoded SSIM 0.7403 vs eager per the PR body).

Head cbab605ef already carries the review's F1 blocker fix upstream
(ac98869aa: int64 row offsets in the fused qknorm+RoPE kernel - verified
present at qknorm_rope.py:43, tl.program_id(0).to(tl.int64)) plus the
engagement-test/logging hardening (cbab605ef), so no fix needs to be carried
by this stack for #1735.

Clean merge, no conflicts. Note: #1732 does not touch minimax_h3.py or envs.py
(its true merge-base is 0462e1b0e; earlier apparent overlap was #1712/#1290
noise from diffing against the wrong base), so 1735's DiT + envs.py edits had
no sibling edits to reconcile.
2026-08-21 10:25:09 +00:00
Will Lin 160f0c9ccf Merge PR #1732 head (ac56806af) into integration/h3-perf-stack
[perf] Optimize MiniMax-H3 text encoder memory: slim single-tensor forward
contract for the Qwen3-VL conditioner (early return at the layer-50 tap,
-262 MiB peak on top of merged #1711), TextEncoder base relaxed to
Generic[TextEncoderOutputT], and opt-in checkpoint-serialized block-FP8
(new text_encoder_quantization.py + minimax_h3_checkpoint_fp8.py).

Clean merge, no conflicts (PR base c4ad4227c == main's state for all touched
files; the tests/local_tests/minimax_h3/README.md hunk applied cleanly on top
of the #1703 dedupe).

Review status (pr1732.md): the two blockers are sm12x/FP8-only - the cutlass
m%4 crash is on the sm12x route (GB200/sm100 uses trtllm, verified clean at
m=559), and the FP8+cpu-offload gap only matters with FP8 enabled. This stack
keeps text-encoder FP8 OFF for all benches, so neither blocker is reachable.
2026-08-21 10:22:26 +00:00
Will Lin e5d1110a0f Merge PR #1703 head (pr1703-fixes @ 942f7db3d) into integration/h3-perf-stack
#1703 (H3 VAE peak-memory streaming) was already squash-merged into main as
e0a3db565 INCLUDING the pr1703-fixes review commits (aadb23f40 per-plane async
pinned copies, 942f7db3d legacy-decode oracle + slicing + pinned-buffer tests) -
the fix-branch tip is byte-identical to main for minimax_h3_video.py, the H3
stages, and the streaming tests. Content no-op recording the reviewed head.

Resolution: the auto-merge textually duplicated the 'Video VAE memory benchmark'
section in tests/local_tests/minimax_h3/README.md (main's squash placed the same
block at a slightly different anchor); deduplicated to a single copy - final
README is byte-identical to origin/main's.
2026-08-21 10:22:00 +00:00
Will Lin 622217ff2a Merge PR #1362 head (pr1362-fixes @ aa95a4c18) into integration/h3-perf-stack
#1362 (on-device uint8 post-decode) was already squash-merged into main as
fca45bc8e INCLUDING the pr1362-fixes review commits (15a164a05 size-regression
fix, f56f56704 comment accuracy, aa95a4c18 CPU regression tests) - the fix-branch
tip is byte-identical to main for video_generator.py and its tests. This merge is
therefore a content no-op recording the reviewed head in the stack history.
No conflicts.
2026-08-21 10:20:46 +00:00
Will LinandClaude Fable 5 cbab605eff [misc] MiniMax-H3 fusions: exact eager fallback, enable/inert logging, engagement test
- _can_run_minimax_h3_fusion now also requires Triton availability, so an
  enabled fusion on a CUDA build without a working Triton falls back to
  eager instead of hitting the strict wrappers' RuntimeError mid-forward.
- One-time logger.info of the resolved fusion set at model init (and a
  warning when the set is requested without Triton), plus a
  prepare_for_compile hook warning that torch.compile capture makes the
  fusions inert inside compiled block forwards.
- Positive-engagement routing test: counts 1/1/2/1 fused-kernel calls in
  one CUDA inference block forward and asserts a grad-enabled forward
  leaves the counters unchanged, so a guard regression to always-eager
  can no longer pass silently.
- Document the index-bounds contract ([0, table_rows)) on the modulation
  wrappers; a device-side check would synchronize, and in-model callers
  are safe by construction.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 10:06:23 +00:00
Will LinandClaude Fable 5 ac98869aa1 [bugfix] MiniMax-H3 fusions: int64 row offsets in the fused qknorm+RoPE kernel
_qknorm_partial_rope_kernel left tl.program_id(0) in int32, so
row * head_dim wrapped once the flattened input reached 2**31 elements
and the kernel read/wrote out of bounds (CUDA illegal memory access).
The PR's other two kernels (modulation.py, swiglu.py) already cast
tl.program_id(0).to(tl.int64); this one now matches, and seq_index /
table_offset inherit int64 from row.

Confirmed on GB200: (1, 8_500_000, 2, 128) bf16 (2.176e9 elements)
crashed before the cast and matches eager after it; the just-under-2**31
control shape matched all along. For H3 (56 heads x 128 head_dim) the
boundary is batch*seq >= 299_593 tokens per rank, reachable at SP=1.
Adds a GPU regression test at the over-2**31 shape that compares the
head and tail rows against eager.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 10:06:12 +00:00
William LinandClaude Fable 5 56d4a6074f [bugfix] fastvideo-kernel: fix Triton block-sparse backward logit scaling (bf16 K pre-scaling) (#1730)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 04:50:06 -05:00
Will LinandClaude Fable 5 9713ea1275 [misc] use the full release name FastVideo-Minimax-FastH3-Preview-v0.1
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Will LinandClaude Fable 5 37aa382cce [bugfix] example: the distilled FastH3 grid is 4 forwards = a 5-point sigma grid
MiniMaxH3Scheduler.set_timesteps(N) builds an N-point sigma grid ending at
0 and runs N-1 transformer forwards (the base model's '50 steps' preset is
49 forwards). The student was distilled on a 4-FORWARD grid
(t = 1000/750/500/250 -> 0 on the shift-12 schedule), so --steps 4 was
silently running a 3-forward, off-distribution grid. Default is now 5
grid points = the distilled 4-forward grid, with the convention documented
on the flag.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Will LinandClaude Fable 5 3f00983287 [feat] attention: opt-in sm_100a CUDA forward route for VSA-H3 tile-64
FASTVIDEO_VSA_SM100A=1 (default off) sends no-grad tile-64 forwards through
fastvideo_kernel.block_sparse_attn_sm100a (the Blackwell block-sparse
forward merged in #1719, which handles per-q-tile NON-uniform q2k_num rows
and zero-count rows). Route preconditions: module importable,
is_supported() (sm_100 device, bf16, head_dim 128, even tile count), no
grad tracking; any failure with the env set logs one warning and falls
back to the Triton-64 kernels, which also keep the entire grad path
unchanged.

The bool block map is compacted with the same map_to_index the Triton bool
entry uses, so the sm_100a kernel sees H3's true NON-uniform per-row counts
(prefix query tiles dense, video tiles prefix+top-k).

The FastH3 example surfaces the route as --vsa-kernel {triton,sm100a}
(default triton), which sets the env before pipeline boot so spawned GPU
workers inherit it; documented in the basic README.

CPU route-selection tests: default-off, engage-on-env, grad fallback,
warn-once fallback for missing module/unsupported geometry, no-grad-context
detection.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Will LinandClaude Fable 5 2dc57f4070 [misc] rebrand the few-step preview to FastH3 (model string, example name, docs)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Will LinandClaude Fable 5 9df19be719 [feat] example: few-step FastVideo-Minimax-H3-Preview inference (4-step DMD2 student)
basic_fast_minimax_h3.py runs FastVideo/FastVideo-Minimax-H3-Preview-v0.1,
the data-free-DMD2 distillation of MiniMax-H3: 4 denoising steps on the
release sampler's shift-12 schedule (vs the base model's 50), synchronized
video+audio in one pipeline call, guidance_scale 1.0.

The script always requests the VSA-H3 attention backend through the typed
boot-time route (pipeline.experimental -> FastVideoArgs.attention_backend):
the student checkpoint carries trained to_gate_compress gates, which only
exist under that backend. --vsa-sparsity defaults to 0.0 (every tile
selected — exactly dense attention); --vsa-tile-size defaults to 64, the
geometry the student was trained with, and is forwarded even at sparsity 0
because the gate-compress branch pools per tile. The HF repo is private
while the MiniMax H3 Community License review completes; --model-path
accepts a local snapshot meanwhile (noted in the script and README).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Will LinandClaude Fable 5 089eea3970 [feat] inference: expose the VSA-H3 tile size on the run-level route
FastVideoArgs gains VSA_tile_size (default 256, CLI --VSA-tile-size). It
rides the same boot-time route as run-level sparsity
(pipeline.experimental -> FastVideoArgs) and the H3 denoising stage
forwards it to MiniMaxH3VSAMetadataBuilder.build(tile_size=...), which
validates the value against VSA_H3_TILE_SHAPES. 256 keeps today's
behavior everywhere; 64 selects the native 64-token Triton block-sparse
path (FASTVIDEO_VSA_CUTEDSL does not apply there). Only the H3 stage
consumes it; Wan/LTX-2 VSA paths are untouched.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Will LinandClaude Fable 5 dd8447ecc5 [feat] VSA-H3: 64-token tile option on the native Triton block-sparse path
The tile size becomes selectable at metadata build time:
MiniMaxH3VSAMetadataBuilder.build(tile_size=...) accepts 256 (default,
unchanged (4,8,8) tiles and VSA-256 CuTe/Triton routing) or 64. At 64 the
tiles are (4,4,4) and the block map is already at the Triton kernels'
native granularity, so forward and backward run
fastvideo_kernel.block_sparse_attn directly (BHSD, transposed around the
call like the 256 wrapper's Triton branch) with no 256->64 mask expansion;
FASTVIDEO_VSA_CUTEDSL does not apply at 64. Tile geometry, pooled scoring,
prefix chunking, the probe, and the gate_compress views all follow the
configured tile element count.

Also adds _validate_h3_tile_geometry: a synchronous, lru-cache-scoped
bounds check on every built geometry (per-tile sizes in (0, tile_elems],
sizes sum to the packed length, untile index injective into non-pad
slots). Malformed geometry now raises at build time with the numbers in
hand instead of surfacing as an unattributable async device fault at some
later kernel or collective.

CPU tests: hand-computed (4,4,4) oracle on a grid ragged in all three
dims, the production packed shape (768x1344, 124 frames) under both tile
sizes, sparsity-0 SDPA equivalence at tile 64, the guard's 64-bound, and
builder rejection of unknown tile sizes. GPU parity of the 64 route was
validated out of band on GB200: sparsity-0 forward+input-grad parity vs
dense SDPA 3.1e-3 rel-L2 at the production packed shape (matching the 256
route to the third digit); sparsity-0.9 out/dq bitwise same-seed
deterministic, dk/dv ~3e-5 reduction-order drift (same profile as the
existing 256 Triton route).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 09:27:38 +00:00
Davids048 dca423fd31 Enable H3 VAE attention backend selection. 2026-08-21 08:33:09 +00:00
Davids048 1b43af8e8e Custom compile H3 vae decoding. 2026-08-21 08:33:09 +00:00
Davids048 628591b620 Add some logging. 2026-08-21 08:33:09 +00:00
Davids048 0980ca563f Add nvtx profiling support. 2026-08-21 08:33:09 +00:00
H1yori233 b158388733 [perf] Add opt-in MiniMax-H3 Sol-Engine fusions 2026-08-21 00:20:38 -07:00
Shao Duan c4ad4227c0 [misc] MiniMax-H3: move the AdaLN converter into scripts/checkpoint_conversion (#1712) 2026-08-21 02:15:00 -05:00
KyleNeverGivesUpandSolitaryThinker a63ccce73d [docs]: add a maintained inference cookbook (#1290)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-08-21 01:53:44 -05:00
H1yori233andWill Lin ac56806aff [perf] Optimize MiniMax-H3 text encoder
Co-authored-by: Will Lin <160547796+KyleNeverGivesUp@users.noreply.github.com>
2026-08-20 21:42:53 -07:00
KyleNeverGivesUp 0462e1b0e7 [perf]: MiniMax H3 - build the Qwen3-VL encoder only as far as it is read (-13.7 GB) (#1711) 2026-08-20 23:19:25 -05:00
lpc0220 907f2100ec [kernel] sm_100a CUDA block-sparse VSA forward (Blackwell), 64- and 128-token blocks (#1719) 2026-08-20 23:16:25 -05:00
Kaiqin Kong e0a3db5651 [perf] Reduce MiniMax-H3 VAE peak memory (#1703) 2026-08-20 23:13:23 -05:00
Raghav K fca45bc8e1 [perf] Quantize frames to uint8 on-device before the post-decode D->H copy (#1362) 2026-08-20 21:55:55 -05:00
Will LinandClaude Fable 5 aa95a4c18e [misc]: add CPU regression tests for the on-device uint8 frame path
Two tests certifying #1362's frame semantics without a GPU:

- frames_match_legacy_cpu_loop: the quantize-then-grid path reproduces
  the legacy make_grid -> *255 -> uint8 per-frame loop bit-exactly for
  in-range fp32 pixels (batch>1 nrow=6 grid layout, odd frame count,
  uint8 HWC contract). Verified to pass against main's legacy loop too,
  so it pins both sides of the equivalence.
- frames_clamp_out_of_range_pixels: out-of-[0,1] VAE output saturates
  at 0/255 instead of wrapping mod 256; fails on the pre-#1362 loop.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 02:39:09 +00:00
Will LinandClaude Fable 5 942f7db3db [test] H3 VAE streaming: legacy-decode oracle, batch slicing, pinned-buffer coverage
Pin the chunk-iterator refactor to the pre-streaming _decode implementation
bit-for-bit across the seam and pad-trim geometries (one padded chunk, a pad
hitting the intra-clip tail, two blended chunks, three chunks plus trim), and
cover the use_slicing batch paths for encode_pixels/decode_to_pixels and the
CUDA pinned-buffer async copy path (skipped without a GPU).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 02:38:44 +00:00
Will LinandClaude Fable 5 aadb23f409 [perf] H3 VAE streamed decode: direct per-plane copies, async with pinned buffers
The temporal slice of the CPU output buffer is strided across channels, so
each finalized-chunk copy_ staged through a pageable CPU temporary plus a
CPU-side scatter, which is where the streamed path's decode-time regression
came from and why the pinned buffer bought nothing. Copy per (batch, channel)
plane instead - contiguous on both sides, memcpy-eligible - and make the
copies non_blocking when the destination is pinned, draining the stream once
in decode_to_pixels before the buffer can be read or released.

Also: raise on a non-positive decode plan before allocating the output
buffer, annotate _decode_chunks as an Iterator, and document the
encode_pixels CPU dtype/range contract.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 02:38:35 +00:00
Will LinandClaude Fable 5 f56f567042 [misc]: correct the post-decode rationale comments
The pre-#1362 `samples.copy_(output_batch.output)` never passed
`non_blocking=True` (git log -S confirms), so drop the deferred
non-blocking-transfer claim and describe the measured costs instead:
a full fp32 D->H copy plus a single-threaded per-frame CPU loop. Also
drop the reintroduced hardcoded "~50 MB" size estimate (same class of
comment Copilot flagged and commit 0399713e7 removed elsewhere) and
scope the "typical flow" claim to the CLI, since the SamplingParam
API default is return_frames=True.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 02:34:43 +00:00
Will LinandClaude Fable 5 15a164a052 [bugfix]: derive result size from the decoded output when the samples mirror is skipped
PR #1362 commit 3 leaves `samples` as an empty placeholder when
`return_frames=False`, but main picked up #1595's refiner size
reporting in the meantime, and `_resolve_output_size(samples, ...)`
silently fell back to the requested geometry in exactly the common
save flow the PR optimizes. Read the geometry from
`output_batch.output` instead (shape-only access, no D->H copy),
gated on `needs_frame_output` so metadata-only and audio-only calls
still never inspect the (possibly dropped) worker output.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 02:34:16 +00:00
Will Lin aaaa7a14a3 Merge remote-tracking branch 'origin/main' into pr1362-fixes 2026-08-21 02:21:44 +00:00
Aryan Kumar 86d639c848 [docs] Announce FastMetal-QAD (#1721) 2026-08-19 13:45:49 -07:00
00338aa9ca [perf] Add FA4 CuTe backward support for VSA-256 (#1639)
Co-authored-by: Hyunsung Lee <hyunsungl@sizigistudios.com>
Co-authored-by: alexzms <3036648523@qq.com>
2026-08-19 11:46:49 -07:00
H1yori233 74b409d7cf [test]: add H3 VAE parity and memory benchmark 2026-08-19 02:45:23 -07:00
Aryan KumarandAryan Kumar 8537dcd6de [feat]: Apple Silicon MLX runtime — INT8 Wan2.1 and Wan2.2 inference (#1638)
Co-authored-by: Aryan Kumar <aryan5v@users.noreply.github.com>
2026-08-18 15:08:00 -07:00
H1yori233 528cef02c4 optimize VAE memory 2026-08-12 02:48:43 -07:00
William Lin 8208536cd1 [bugfix] profiler region system: record + export actually work, usability roll-up (#1691) 2026-08-09 20:35:53 -07:00
Kai 0653f8f3af [new-model] Add V2A: native MMAudio inference pipeline (#1622) 2026-08-09 17:29:27 -07:00
William Lin e0d702decb [feat] VSA for MiniMax H3: packed mixed-modality sparse attention (#1695) 2026-08-09 13:10:51 -07:00
Shao Duan 541ef014ee [perf] MiniMax-H3: rank-reduced AdaLN pruned model option (-39% params, -23 GiB VRAM) (#1699) 2026-08-09 12:31:57 -07:00
William Lin ffc1a7a58b [refactor] H3 pipeline cleanup: shared helpers, dead machinery, loop-invariant hoists (#1698) 2026-08-09 04:51:31 -07:00
William Lin 9028953625 [misc] yapf pass under CI's interpreter (3.12) + pin hook language_version (#1702) 2026-08-08 22:24:17 -07:00
KyleNeverGivesUpandClaude Opus 5 6eb95693a1 [misc]: re-run yapf on main so pre-commit passes again (#1700)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-07 16:17:52 -07:00
Junda Su c3567eb468 [feat] add Minimax H3 sft pipeline (#1688) 2026-08-07 16:12:04 -07:00
RaghavandClaude Opus 4.7 3c3da4d057 [perf]: skip the fp32 samples D->H copy when return_frames=False
The previous commit rebuilt the post-decode frames path to read
`output_batch.output` directly via the GPU `vid_u8` cast, so `samples`
is now consumed in exactly one place — the result dict's `samples`
field, gated on `batch.return_frames`. When the caller doesn't ask
for `samples`, the pinned ~50 MB fp32 alloc and its D->H copy (and
the latent fallthrough `.cpu()`) are dead weight; the typical
generate flow (save_video=True, return_frames=False) hits this on
every call.

Extend `skip_pixel_prealloc` to include `not return_frames` so the
pinned buffer is allocated only when needed, and short-circuit the
copy/`.cpu()` on the same condition. No effect when
`return_frames=True` or for latent callers that read `samples`;
correctness is unchanged (SSIM-gated, same as the parent commit).

Removes the residual ~2 s the PR text already calls out for the
"only saving to disk" case.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-16 03:02:49 -07:00
Raghav 0399713e7b [perf]: clarify samples/output_batch.output equivalence; drop hardcoded transfer size
Address review comments on #1362:
- Document that `samples` is just the pinned-CPU mirror of
  `output_batch.output` with no intervening preprocessing, so sourcing
  from `output_batch.output` is the same data (Copilot).
- Replace the hardcoded ~0.3-1.4 GB estimate with a description of the
  scaling relationship; it varies with resolution/frames/batch/dtype
  and would rot (Copilot).
- Note the SSIM-gated (not bit-exact) equivalence inline.

No behavior change.
2026-07-16 02:59:59 -07:00
Raghav 0af2e9e8ef [perf]: quantize frames to uint8 on-device before the D->H copy
PostDecodeFrameProcessStage was charged ~6s/run (~25% of e2e on Cosmos
2.5, scaling with frames/resolution). Profiling (nsys + microbench)
showed the stage's own compute is only ~0.5-1.9s; the rest is the
non-blocking pinned-CPU samples.copy_(output) D->H of the full fp32
video (~0.3-1.4 GB) completing lazily and blocking the first
postprocess op, plus a single-threaded per-frame CPU *255/cast loop.

Cast to uint8 on the source device (typically CUDA) before the copy:
the transfer becomes 4x smaller (fp32 -> uint8) and the elementwise
work runs on the GPU. Microbench: 2.056s -> 0.034s (T=29), 2.304s ->
0.132s (T=125), 17-60x on the measurable cost.

clamp_(0, 255) additionally fixes a latent overflow: VAE output
slightly outside [0, 1] previously wrapped mod 256 in the unclamped
(x * 255).to(uint8) cast.

Output is not bit-identical to the old CPU cast (float->uint8 differs
<=1 LSB between CPU and GPU on boundary pixels); gate via SSIM rather
than exact equality. Scoped to the pixel-video path only; latent,
audio-only, and return-samples paths are unchanged.
2026-07-16 02:59:59 -07:00
806 changed files with 52958 additions and 25290 deletions
@@ -10,9 +10,7 @@ from pathlib import Path
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Clone a reference repo for FastVideo parity tests."
)
parser = argparse.ArgumentParser(description="Clone a reference repo for FastVideo parity tests.")
parser.add_argument("repo_url", help="Official reference repository URL")
parser.add_argument("target_dir", help="Directory to clone into")
parser.add_argument("--branch", help="Branch or tag to clone")
@@ -62,9 +60,7 @@ def gitignore_entry_for(target: Path) -> str:
try:
relative = resolved.relative_to(root)
except ValueError as exc:
raise ValueError(
"--update-gitignore requires target_dir to be under the current directory"
) from exc
raise ValueError("--update-gitignore requires target_dir to be under the current directory") from exc
text = relative.as_posix().rstrip("/")
return "/" + text + "/"
@@ -8,14 +8,12 @@ import os
import sys
from pathlib import Path
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Download a HF model snapshot or selected files into a local directory."
)
description="Download a HF model snapshot or selected files into a local directory.")
parser.add_argument("repo_id", help="HF repo id, for example Org/Model")
parser.add_argument("local_dir", help="Destination directory")
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
@@ -10,7 +10,6 @@ import sys
from pathlib import Path
from typing import Any
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
RAW_WEIGHT_SUFFIXES = (".safetensors", ".pt", ".pth", ".ckpt", ".bin")
KNOWN_COMPONENTS = {
@@ -34,8 +33,7 @@ KNOWN_COMPONENTS = {
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown."
)
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown.")
parser.add_argument("source", help="HF repo id or local weights directory")
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
parser.add_argument("--revision", help="HF revision to inspect")
@@ -94,14 +92,12 @@ def load_remote_files(
) -> list[str]:
from huggingface_hub import list_repo_files
return sorted(
list_repo_files(
repo_id,
repo_type=repo_type,
revision=revision,
token=token,
)
)
return sorted(list_repo_files(
repo_id,
repo_type=repo_type,
revision=revision,
token=token,
))
def load_remote_model_index(
@@ -215,24 +211,24 @@ def build_result(args: argparse.Namespace) -> dict[str, Any]:
"components_seen": components,
"file_count": len(files),
"file_scan_truncated": truncated,
"files_sample": files[: args.sample_limit],
"files_sample": files[:args.sample_limit],
}
def print_human(result: dict[str, Any]) -> None:
for key in (
"source",
"source_kind",
"repo_type",
"revision",
"token_env",
"source_layout",
"needs_conversion",
"model_index_class",
"model_index_diffusers_version",
"model_index_error",
"file_count",
"file_scan_truncated",
"source",
"source_kind",
"repo_type",
"revision",
"token_env",
"source_layout",
"needs_conversion",
"model_index_class",
"model_index_diffusers_version",
"model_index_error",
"file_count",
"file_scan_truncated",
):
value = result.get(key)
if value is not None:
@@ -18,7 +18,6 @@ import pytest
import torch
from torch.testing import assert_close
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29519")
os.environ.setdefault("DISABLE_SP", "1")
@@ -35,15 +34,10 @@ FASTVIDEO_CONFIG_CLASS = "<FastVideoConfig>" # TODO.
FASTVIDEO_MODEL_MODULE = "fastvideo.models.<bucket>.<module>" # TODO.
FASTVIDEO_MODEL_CLASS = "<FastVideoModel>" # TODO.
OFFICIAL_REF_DIR = Path(
os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>")
)
LOCAL_WEIGHTS_DIR = Path(
os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY)
)
CONVERTED_WEIGHTS_DIR = Path(
os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / FAMILY)
)
OFFICIAL_REF_DIR = Path(os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>"))
LOCAL_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY))
CONVERTED_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR",
REPO_ROOT / "converted_weights" / FAMILY))
def _resolve_hf_token() -> str | None:
@@ -99,18 +93,14 @@ def _load_official_model(device: torch.device, dtype: torch.dtype) -> torch.nn.M
model = OfficialClass() # TODO: pass official config kwargs.
state_dict = {} # TODO: load official state dict from LOCAL_WEIGHTS_DIR.
missing, unexpected = model.load_state_dict(state_dict, strict=True)
assert not missing and not unexpected, (
f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
)
assert not missing and not unexpected, (f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
"""Load the FastVideo component with the same tensor content."""
if not CONVERTED_WEIGHTS_DIR.exists() and not LOCAL_WEIGHTS_DIR.exists():
pytest.skip(
f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}"
)
pytest.skip(f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}")
# TODO: replace with the bucket-specific FastVideo config/class/loader.
# DiT examples:
@@ -127,8 +117,7 @@ def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.
state_dict = {} # TODO: load converted or directly mapped state dict.
missing, unexpected = model.load_state_dict(state_dict, strict=True)
assert not missing and not unexpected, (
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
)
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
return model.to(device=device, dtype=dtype).eval()
@@ -187,11 +176,9 @@ def test_component_parity():
assert official_out.shape == fastvideo_out.shape
diff = (official_out - fastvideo_out).abs()
print(
f"official abs_mean={official_out.abs().mean().item():.6f} "
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}"
)
print(f"official abs_mean={official_out.abs().mean().item():.6f} "
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
# TODO: pick tolerance by scope:
# - single block / same kernel: 1e-4
@@ -27,7 +27,6 @@ try:
except ImportError: # pragma: no cover - optional local conversion dependency
snapshot_download = None
# TODO: fill with authoritative component prefixes for monolithic checkpoints.
# Example: {"model.model.": "transformer", "pretransform.model.": "vae"}
COMPONENT_PREFIXES: dict[str, str] = {}
@@ -47,10 +46,7 @@ SKIP_PATTERNS: tuple[str, ...] = ()
def _hf_token() -> str | None:
return (
os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
or os.environ.get("HF_API_KEY")
)
return (os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN") or os.environ.get("HF_API_KEY"))
def resolve_src(src: str, revision: str | None) -> Path:
@@ -95,11 +91,10 @@ def apply_mapping(key: str) -> str | None:
return key
def split_monolithic(
state: dict[str, torch.Tensor],
) -> dict[str, OrderedDict[str, torch.Tensor]]:
def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[str, torch.Tensor]]:
components: dict[str, OrderedDict[str, torch.Tensor]] = {
name: OrderedDict() for name in set(COMPONENT_PREFIXES.values())
name: OrderedDict()
for name in set(COMPONENT_PREFIXES.values())
}
intentionally_skipped: list[str] = []
unowned: list[str] = []
@@ -117,10 +112,8 @@ def split_monolithic(
unowned.append(key)
if unowned:
sample = ", ".join(unowned[:10])
raise ValueError(
f"Unowned monolithic keys: {len(unowned)}. "
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}"
)
raise ValueError(f"Unowned monolithic keys: {len(unowned)}. "
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}")
if intentionally_skipped:
print(f"Intentionally skipped {len(intentionally_skipped)} keys")
return {name: weights for name, weights in components.items() if weights}
@@ -143,8 +136,12 @@ def build_component_configs(_src_dir: Path) -> dict[str, dict[str, Any]]:
# TODO: emit config content accepted by FastVideo loaders. Most components use
# config.json; schedulers use scheduler_config.json.
return {
"transformer": {"_class_name": "<FastVideoTransformerClass>"},
"vae": {"_class_name": "<FastVideoVAEClass>"},
"transformer": {
"_class_name": "<FastVideoTransformerClass>"
},
"vae": {
"_class_name": "<FastVideoVAEClass>"
},
}
@@ -177,19 +174,13 @@ def build_model_index(
}
if revision:
index["_fastvideo_converted_revision"] = revision
return {
key: value
for key, value in index.items()
if key.startswith("_") or key in available_components
}
return {key: value for key, value in index.items() if key.startswith("_") or key in available_components}
def validate_component_configs(configs: dict[str, dict[str, Any]]) -> None:
# TODO: instantiate each FastVideo config and call update_model_arch(...) or
# update_model_config(...) with this JSON so unknown emitted keys fail here.
placeholder_configs = [
name for name, config in configs.items() if "<" in json.dumps(config)
]
placeholder_configs = [name for name, config in configs.items() if "<" in json.dumps(config)]
if placeholder_configs:
raise ValueError(f"Replace config placeholders for: {placeholder_configs}")
@@ -201,9 +192,7 @@ def verify_conversion(
del dst_dir, components
# TODO: load each emitted stateful component through its production loader and
# assert strict load, or document exact allowed missing/unexpected keys.
raise NotImplementedError(
"Implement production config validation and strict-load checks"
)
raise NotImplementedError("Implement production config validation and strict-load checks")
def write_component(
@@ -216,9 +205,7 @@ def write_component(
if component_dir.exists() and any(component_dir.iterdir()):
shutil.rmtree(component_dir)
component_dir.mkdir(parents=True, exist_ok=True)
save_file(
dict(state), str(component_dir / "diffusion_pytorch_model.safetensors")
)
save_file(dict(state), str(component_dir / "diffusion_pytorch_model.safetensors"))
if config is not None:
config_path = component_dir / config_filename(name)
with config_path.open("w", encoding="utf-8") as f:
@@ -261,9 +248,7 @@ def convert(
if layout in {"monolithic", "raw_official"}:
# TODO: replace model.safetensors with the official monolithic file name.
components = split_monolithic(
load_checkpoint(default_monolithic_checkpoint(src_path))
)
components = split_monolithic(load_checkpoint(default_monolithic_checkpoint(src_path)))
elif layout in {"separate_components", "mixed"}:
if not src_path.is_dir():
raise ValueError(f"{layout} layout requires a source directory: {src_path}")
@@ -271,9 +256,7 @@ def convert(
else:
raise ValueError(f"Unsupported template layout: {layout}")
copied = (
copy_passthrough(src_path, dst_dir) if src_path.is_dir() else []
)
copied = (copy_passthrough(src_path, dst_dir) if src_path.is_dir() else [])
configs = build_component_configs(src_path if src_path.is_dir() else src_path.parent)
validate_component_configs(configs)
for name, state in components.items():
@@ -289,9 +272,7 @@ def convert(
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--src", required=True, help="HF repo id, local dir, or checkpoint path"
)
parser.add_argument("--src", required=True, help="HF repo id, local dir, or checkpoint path")
parser.add_argument("--revision", help="HF branch, tag, or commit for repo sources")
parser.add_argument(
"--dst",
@@ -24,8 +24,8 @@ from typing import Any
import torch
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
DRILL_LAYER_ENV: str = "<FAMILY>_DEBUG_DRILL_LAYER"
HYPOTHESIS_ENV: str = "<FAMILY>_DEBUG_PATCH_<HYPOTHESIS>"
REL_THRESHOLD: float = 0.005 # 0.5% abs_mean drift flags a block as divergent
@@ -94,6 +94,7 @@ def _attach_block_hooks(
handles: list[Any] = []
def _hook(name: str):
def fn(_module, _inputs, outputs):
t = outputs[0] if isinstance(outputs, tuple) else outputs
if not torch.is_tensor(t):
@@ -101,6 +102,7 @@ def _attach_block_hooks(
log.append({"side": label, **_stat(name, t)})
if tensors is not None:
tensors[name] = t.detach().float().cpu()
return fn
def _pre_hook(name: str):
@@ -114,6 +116,7 @@ def _attach_block_hooks(
log.append({"side": label, **_stat(key, t)})
if tensors is not None:
tensors[key] = t.detach().float().cpu()
return fn
# TODO: adapt attribute paths to your model. Remove adapter block if absent.
@@ -131,43 +134,21 @@ def _attach_block_hooks(
# magi-human uses: attention, mlp.pre_norm, mlp.up_gate_proj,
# mlp.down_proj (pre+post), mlp, attn_post_norm, mlp_post_norm.
if hasattr(layer, "attention"):
handles.append(
layer.attention.register_forward_hook(_hook(f"{tag}.attention"))
)
handles.append(layer.attention.register_forward_hook(_hook(f"{tag}.attention")))
if hasattr(layer, "mlp"):
mlp = layer.mlp
if hasattr(mlp, "pre_norm"):
handles.append(
mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm"))
)
handles.append(mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm")))
if hasattr(mlp, "up_gate_proj"):
handles.append(
mlp.up_gate_proj.register_forward_hook(
_hook(f"{tag}.mlp.up_gate_proj")
)
)
handles.append(mlp.up_gate_proj.register_forward_hook(_hook(f"{tag}.mlp.up_gate_proj")))
if hasattr(mlp, "down_proj"):
handles.append(
mlp.down_proj.register_forward_pre_hook(
_pre_hook(f"{tag}.mlp.down_proj")
)
)
handles.append(
mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj"))
)
handles.append(mlp.down_proj.register_forward_pre_hook(_pre_hook(f"{tag}.mlp.down_proj")))
handles.append(mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj")))
handles.append(mlp.register_forward_hook(_hook(f"{tag}.mlp")))
if hasattr(layer, "attn_post_norm"):
handles.append(
layer.attn_post_norm.register_forward_hook(
_hook(f"{tag}.attn_post_norm")
)
)
handles.append(layer.attn_post_norm.register_forward_hook(_hook(f"{tag}.attn_post_norm")))
if hasattr(layer, "mlp_post_norm"):
handles.append(
layer.mlp_post_norm.register_forward_hook(
_hook(f"{tag}.mlp_post_norm")
)
)
handles.append(layer.mlp_post_norm.register_forward_hook(_hook(f"{tag}.mlp_post_norm")))
return handles
@@ -193,11 +174,9 @@ def _write_log(entries: list[dict], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w") as f:
for e in entries:
f.write(
f"{e['name']} {e['shape']} "
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
f"{e['min']:.6f} {e['max']:.6f}\n"
)
f.write(f"{e['name']} {e['shape']} "
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
f"{e['min']:.6f} {e['max']:.6f}\n")
def _sort_key(name: str, drill_layer: int) -> tuple:
@@ -205,9 +184,14 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
return (0, "")
if name.startswith(f"L{drill_layer:02d}."):
sub_order = {
"attention": 0, "attn_post_norm": 1, "mlp.pre_norm": 2,
"mlp.up_gate_proj": 3, "mlp.down_proj<in>": 4,
"mlp.down_proj": 5, "mlp": 6, "mlp_post_norm": 7,
"attention": 0,
"attn_post_norm": 1,
"mlp.pre_norm": 2,
"mlp.up_gate_proj": 3,
"mlp.down_proj<in>": 4,
"mlp.down_proj": 5,
"mlp": 6,
"mlp_post_norm": 7,
}.get(name.split(".", 1)[1], 9)
return (1, f"block[{drill_layer:02d}]", sub_order)
if name.startswith("block["):
@@ -216,10 +200,8 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
hdr = (
f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}"
)
hdr = (f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}")
print(f"\n{hdr}\n{'-' * len(hdr)}")
first_div: int | None = None
for name in sorted(by_name.keys(), key=lambda n: _sort_key(n, drill_layer)):
@@ -235,11 +217,9 @@ def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
flag = " <<< DIVERGE"
if first_div is None:
first_div = int(name[len("block["):-1])
print(
f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}"
)
print(f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}")
return first_div
@@ -255,10 +235,8 @@ def _print_elementwise(up_t: dict[str, torch.Tensor], fv_t: dict[str, torch.Tens
continue
diff = (a - b).abs()
rel = (diff.mean().item() / max(a.abs().mean().item(), 1e-9)) * 100
print(
f"{name:<30} {str(tuple(a.shape)):<22} "
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%"
)
print(f"{name:<30} {str(tuple(a.shape)):<22} "
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%")
def main() -> None:
@@ -43,12 +43,10 @@ def _add_official_to_path() -> Path:
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
value = tensor.detach().float()
print(
f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
f"dtype={tensor.dtype} device={tensor.device} "
f"min={value.min().item():.6f} max={value.max().item():.6f} "
f"mean={value.mean().item():.6f} std={value.std().item():.6f}"
)
print(f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
f"dtype={tensor.dtype} device={tensor.device} "
f"min={value.min().item():.6f} max={value.max().item():.6f} "
f"mean={value.mean().item():.6f} std={value.std().item():.6f}")
def _extract_tensor(output: Any, key: str) -> torch.Tensor:
@@ -73,10 +71,8 @@ def _run_official_pipeline(
device: torch.device,
) -> Any:
del official_path, params, device
pytest.skip(
"TODO: import the official pipeline/factory, load official weights, "
"run with params, and return the comparison target."
)
pytest.skip("TODO: import the official pipeline/factory, load official weights, "
"run with params, and return the comparison target.")
def _run_fastvideo_pipeline(model_path: Path, params: dict[str, Any]) -> Any:
@@ -146,8 +142,6 @@ def test_todo_model_family_pipeline_official_parity() -> None:
assert official_tensor.shape == fastvideo_tensor.shape
diff = (official_tensor - fastvideo_tensor).abs()
print(
f"diff max={diff.max().item():.6f} "
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
)
print(f"diff max={diff.max().item():.6f} "
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}")
assert_close(fastvideo_tensor, official_tensor, atol=1e-2, rtol=1e-2)
+149
View File
@@ -0,0 +1,149 @@
name: macOS MLX Smoke
on:
pull_request:
branches: [main]
paths:
- ".github/workflows/ci-macos-mlx.yml"
- "fastvideo/mlx_runtime/**"
- "fastvideo/tests/mlx/**"
- "fastvideo/tests/platforms/test_mps_vsa_error.py"
- "fastvideo/platforms/mps.py"
- "fastvideo/platforms/__init__.py"
- "fastvideo/__init__.py"
- "examples/inference/basic/mlx_*.py"
- "fastvideo/benchmarks/mlx_*.py"
- "pyproject.toml"
workflow_dispatch:
permissions:
contents: read
concurrency:
group: macos-mlx-${{ github.ref }}
cancel-in-progress: true
jobs:
mlx-smoke:
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
runs-on: macos-15
timeout-minutes: 25
env:
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
TOKENIZERS_PARALLELISM: "false"
MASTER_ADDR: localhost
MASTER_PORT: "29513"
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
cache: pip
- uses: astral-sh/setup-uv@v3
- name: Install lightweight MLX smoke dependencies
run: |
uv pip install --system \
--index-url https://download.pytorch.org/whl/cpu \
torch==2.11.0 torchvision torchaudio
uv pip install --system \
pytest numpy scipy pillow imageio einops cloudpickle filelock \
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
- name: Show Apple runtime
run: |
python - <<'PY'
import platform
import mlx.core as mx
import torch
print("machine:", platform.machine())
print("processor:", platform.processor())
print("mlx default device:", mx.default_device())
memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable"
print("mlx memory_size:", memory_size)
print("torch:", torch.__version__)
print("torch mps available:", torch.backends.mps.is_available())
PY
- name: Run MLX smoke tests
run: |
python -m pytest \
fastvideo/tests/mlx/test_dmd_sampling.py \
fastvideo/tests/mlx/test_memory_limits.py \
fastvideo/tests/mlx/test_quant_capability.py \
fastvideo/tests/mlx/test_mlx_dit_parity.py \
fastvideo/tests/mlx/test_mlx_compile_parity.py \
fastvideo/tests/mlx/test_mlx_checkpoint.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
fastvideo/tests/mlx/test_mlx_refine.py \
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
fastvideo/tests/mlx/test_wan22_sample.py \
fastvideo/tests/mlx/test_windowed_attention.py \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
fastvideo/tests/platforms/test_mps_vsa_error.py \
-q
# Same tests on MLX's CPU backend. Hosted macOS runners are scarce and
# slower to schedule; this Linux job gives fast PR signal on the identical
# graph (the parity tests were designed to be backend-agnostic), while the
# macOS job above stays the source of truth for Metal behavior.
mlx-smoke-linux-cpu:
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
runs-on: ubuntu-latest
timeout-minutes: 20
env:
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
TOKENIZERS_PARALLELISM: "false"
MASTER_ADDR: localhost
MASTER_PORT: "29513"
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
cache: pip
- uses: astral-sh/setup-uv@v3
- name: Install lightweight MLX smoke dependencies (CPU backend)
run: |
uv pip install --system \
--index-url https://download.pytorch.org/whl/cpu \
torch==2.11.0 torchvision torchaudio
uv pip install --system \
pytest numpy scipy pillow imageio einops cloudpickle filelock \
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
- name: Run MLX smoke tests (CPU backend)
run: |
python -m pytest \
fastvideo/tests/mlx/test_dmd_sampling.py \
fastvideo/tests/mlx/test_memory_limits.py \
fastvideo/tests/mlx/test_quant_capability.py \
fastvideo/tests/mlx/test_mlx_dit_parity.py \
fastvideo/tests/mlx/test_mlx_compile_parity.py \
fastvideo/tests/mlx/test_mlx_checkpoint.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
fastvideo/tests/mlx/test_mlx_refine.py \
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
fastvideo/tests/mlx/test_wan22_sample.py \
fastvideo/tests/mlx/test_windowed_attention.py \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
fastvideo/tests/platforms/test_mps_vsa_error.py \
-q
+2
View File
@@ -6,6 +6,7 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
@@ -16,6 +17,7 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
+1
View File
@@ -23,6 +23,7 @@ Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
outputs_audio/
outputs_video
checkpoints/
sbatch.sh
+1
View File
@@ -22,6 +22,7 @@ repos:
hooks:
- id: yapf
args: [--in-place, --verbose]
language_version: python3.12
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.11.12
+6
View File
@@ -9,6 +9,7 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
@@ -62,6 +63,11 @@ UV_TORCH_BACKEND=cu126 uv pip install fastvideo
Use `UV_TORCH_BACKEND=cu130` on CUDA 13. Apple silicon users should follow the
[MPS installation guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
> **On an Apple Silicon Mac?** FastVideo runs FastWan text-to-video natively
> through an MLX runtime — a 5-second 480p clip generated locally, no cloud,
> no discrete GPU. Install with `uv pip install -e '.[mlx]'` and follow the
> [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
> **On an NVIDIA DGX Spark (GB10 / ARM64 + CUDA 13)?** There's no prebuilt ARM wheel for the FastVideo CUDA kernel, so it's an editable from-source install (`UV_TORCH_BACKEND=cu130 uv pip install -e .`, which compiles that kernel for you) rather than `UV_TORCH_BACKEND=cu130 uv pip install fastvideo`. A compatible prebuilt ARM64 FlashAttention wheel is available separately. Follow the [DGX Spark install guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/spark/).
@@ -3,7 +3,6 @@ from __future__ import annotations
import sys
from pathlib import Path
TESTS_DIR = Path(__file__).resolve().parent
DREAMVERSE_PACKAGE_DIR = TESTS_DIR.parent
DREAMVERSE_APP_DIR = DREAMVERSE_PACKAGE_DIR.parent
@@ -5,7 +5,6 @@ from pathlib import Path
import pytest
SERVER_DIR = Path(__file__).resolve().parents[1]
@@ -53,9 +52,7 @@ def test_config_defaults_to_cerebras_with_parallel_groq_fallback_stage(monkeypat
module = _load_config_module()
assert module.PROMPT_PROVIDER == "cerebras"
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
("cerebras", "groq"),
)
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
assert module.PROMPT_PROVIDER_PRIORITY == (
"cerebras",
"groq",
@@ -86,9 +83,7 @@ def test_config_ignores_legacy_groq_primary_override(monkeypatch):
module = _load_config_module()
assert module.PROMPT_PROVIDER == "cerebras"
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
("cerebras", "groq"),
)
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
assert module.PROMPT_PROVIDER_PRIORITY == (
"cerebras",
"groq",
@@ -106,24 +101,17 @@ def test_config_uses_local_overlay_paths_when_devtools_enabled(monkeypatch, tmp_
assert module.DEVTOOLS_ENABLED is True
assert module.FRONTEND_ROOT.as_posix().endswith("apps/dreamverse/web")
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith(
"dreamverse/prompts.local/next_segment_system_prompt.md"
)
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith("dreamverse/prompts.local/next_segment_system_prompt.md")
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
"dreamverse/prompts/next_segment_system_prompt.md"
)
"dreamverse/prompts/next_segment_system_prompt.md")
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_PATH.endswith(
"dreamverse/prompts.local/rewrite_user_system_prompt.md"
)
"dreamverse/prompts.local/rewrite_user_system_prompt.md")
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
"dreamverse/prompts/rewrite_user_system_prompt.md"
)
"dreamverse/prompts/rewrite_user_system_prompt.md")
assert module.CURATED_PRESETS_FILE_PATH.endswith(
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json"
)
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json")
assert module.CURATED_PRESETS_FALLBACK_FILE_PATH.endswith(
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json"
)
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json")
assert module.FRONTEND_STATIC_DIR_CANDIDATES[:2] == (
str(module.FRONTEND_ROOT / "out"),
str(module.FRONTEND_ROOT / "dist"),
@@ -9,6 +9,7 @@ from fastapi.testclient import TestClient
import fastvideo.entrypoints.streaming as streaming_entrypoints
import pytest
def _install_stack03_import_stubs(monkeypatch):
"""Keep entrypoint tests focused while later-stack runtime modules are absent."""
if not hasattr(streaming_entrypoints, "build_health_router"):
@@ -17,6 +18,7 @@ def _install_stack03_import_stubs(monkeypatch):
gpu_pool_stub = types.ModuleType("dreamverse.gpu_pool")
class GPUPool:
def __init__(self, _gpu_ids):
pass
@@ -49,6 +51,7 @@ def _install_stack03_import_stubs(monkeypatch):
controller_stub = types.ModuleType("dreamverse.session.controller")
class SessionController:
def __init__(self, **_kwargs):
pass
@@ -76,13 +79,11 @@ def _run_cli(module, monkeypatch, argv: list[str]) -> list[dict[str, object]]:
uvicorn_stub = types.ModuleType("uvicorn")
def run(app, host: str, port: int) -> None:
calls.append(
{
"app": app,
"host": host,
"port": port,
}
)
calls.append({
"app": app,
"host": host,
"port": port,
})
uvicorn_stub.run = run
monkeypatch.setitem(sys.modules, "uvicorn", uvicorn_stub)
@@ -99,13 +100,11 @@ def test_server_cli_defaults_to_local_web_port(monkeypatch):
server_main = _import_server_main(monkeypatch)
calls = _run_cli(server_main, monkeypatch, ["dreamverse-server"])
assert calls == [
{
"app": server_main.app,
"host": "0.0.0.0",
"port": 8009,
}
]
assert calls == [{
"app": server_main.app,
"host": "0.0.0.0",
"port": 8009,
}]
def test_server_cli_allows_explicit_host_and_port(monkeypatch):
@@ -116,13 +115,11 @@ def test_server_cli_allows_explicit_host_and_port(monkeypatch):
["dreamverse-server", "--host", "127.0.0.1", "--port", "8123"],
)
assert calls == [
{
"app": server_main.app,
"host": "127.0.0.1",
"port": 8123,
}
]
assert calls == [{
"app": server_main.app,
"host": "127.0.0.1",
"port": 8123,
}]
def test_server_does_not_expose_backend_source_as_static_assets(monkeypatch):
@@ -142,13 +139,11 @@ def test_mock_server_cli_defaults_to_local_web_port(monkeypatch):
["dreamverse-mock-server"],
)
assert calls == [
{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8009,
}
]
assert calls == [{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8009,
}]
def test_mock_server_cli_updates_latency(monkeypatch):
@@ -161,13 +156,11 @@ def test_mock_server_cli_updates_latency(monkeypatch):
["dreamverse-mock-server", "--latency", "321", "--port", "8111"],
)
assert calls == [
{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8111,
}
]
assert calls == [{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8111,
}]
assert mock_server.LATENCY_MS == 321
finally:
mock_server.LATENCY_MS = old_latency_ms
@@ -7,7 +7,6 @@ from types import SimpleNamespace
import pytest
import dreamverse.gpu_pool as gpu_pool
@@ -85,9 +84,7 @@ def test_send_command_raises_on_worker_death():
cmd_q = ctx.Queue()
resp_q = ctx.Queue()
proc = ctx.Process(
target=_child_consume_and_exit, args=(cmd_q, resp_q)
)
proc = ctx.Process(target=_child_consume_and_exit, args=(cmd_q, resp_q))
proc.start()
# Wait for the spawn child to fully boot. Allow generous time —
@@ -7,7 +7,7 @@ ALLOWED_PREFIXES = (
"fastvideo.entrypoints.video_generator",
"fastvideo.configs",
)
ALLOWED_EXACT = ("fastvideo",)
ALLOWED_EXACT = ("fastvideo", )
FORBIDDEN_PREFIXES = (
"fastvideo.pipelines",
"fastvideo.models",
@@ -38,19 +38,13 @@ def test_dreamverse_server_imports_only_public_fastvideo_surfaces() -> None:
except SyntaxError as task_exc:
raise AssertionError(f"Failed to parse {path}") from task_exc
for node in ast.walk(tree):
names = (
[a.name for a in node.names] if isinstance(node, ast.Import)
else [node.module] if isinstance(node, ast.ImportFrom) and node.module
else []
)
names = ([a.name for a in node.names] if isinstance(node, ast.Import) else
[node.module] if isinstance(node, ast.ImportFrom) and node.module else [])
for name in names:
if not name:
continue
rel_path = str(path.relative_to(root))
if (
name.startswith(FORBIDDEN_PREFIXES)
and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS
):
if (name.startswith(FORBIDDEN_PREFIXES) and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS):
bad.append((str(path.relative_to(root)), getattr(node, "lineno", 0), name))
assert bad == [], f"Forbidden internal imports: {bad}"
@@ -6,7 +6,6 @@ import os
from fastapi import WebSocketDisconnect
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -14,6 +13,7 @@ import dreamverse.mock_server as mock_server
class _FakeWebSocket:
def __init__(self, messages: list[tuple[float, dict[str, object]]]):
self._messages = messages
self._index = 0
@@ -49,34 +49,34 @@ def test_mock_server_matches_current_single5s_protocol():
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 1
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "simple_prompt_1",
"curated_prompts": ["selected prompt"],
"single_clip_mode": True,
"enhancement_enabled": False,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.01,
{
"type": "simple_generate",
"preset_id": "simple_custom_prompt",
"prompt_id": "simple_custom_prompt",
"prompt": "custom prompt",
"enhancement_enabled": True,
"initial_image": None,
},
),
(0.20, {"type": "leave"}),
]
)
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "simple_prompt_1",
"curated_prompts": ["selected prompt"],
"single_clip_mode": True,
"enhancement_enabled": False,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.01,
{
"type": "simple_generate",
"preset_id": "simple_custom_prompt",
"prompt_id": "simple_custom_prompt",
"prompt": "custom prompt",
"enhancement_enabled": True,
"initial_image": None,
},
),
(0.20, {
"type": "leave"
}),
])
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -92,24 +92,14 @@ def test_mock_server_matches_current_single5s_protocol():
assert message_types.count("ltx2_stream_complete") == 2
assert "prompt_sources_blocked" not in message_types
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
gpu_assigned_event = next(
payload for payload in ws.sent_json if payload["type"] == "gpu_assigned"
)
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
gpu_assigned_event = next(payload for payload in ws.sent_json if payload["type"] == "gpu_assigned")
assert gpu_assigned_event["session_timeout"] == mock_server.SESSION_TIMEOUT_SECONDS
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
assert segment_start_events[0]["prompt"] == "selected prompt"
assert segment_start_events[1]["prompt"] == "custom prompt"
step_complete_events = [
payload
for payload in ws.sent_json
if payload["type"] == "step_complete"
]
step_complete_events = [payload for payload in ws.sent_json if payload["type"] == "step_complete"]
assert len(step_complete_events) == 2
assert step_complete_events[0]["latency_ms"] == {
"total": 121.0,
@@ -134,29 +124,29 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
mock_server.LATENCY_MS = 1
mock_server.GENERATION_SEGMENT_CAP = 1
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "start a new rollout",
},
),
(0.20, {"type": "leave"}),
]
)
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "start a new rollout",
},
),
(0.20, {
"type": "leave"
}),
])
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -166,11 +156,7 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
assert "generation_cap_reached" not in message_types
assert "prompt_sources_blocked" not in message_types
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
assert segment_start_events[0]["prompt"] == "segment one"
assert segment_start_events[1]["prompt"] == "segment one [start a new rollout]"
@@ -187,54 +173,40 @@ def test_mock_server_rewrite_during_active_segment_restarts_from_first_rewritten
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 100
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one", "segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "restart from rewrite",
},
),
(0.40, {"type": "leave"}),
]
)
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one", "segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "restart from rewrite",
},
),
(0.40, {
"type": "leave"
}),
])
asyncio.run(mock_server.websocket_endpoint(ws))
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
"segment one",
"segment one [restart from rewrite]",
]
assert all(
payload["prompt"] != "segment two"
for payload in segment_start_events[1:]
)
reset_events = [
payload
for payload in ws.sent_json
if payload.get("type") == "seed_prompts_reset_applied"
]
assert any(
payload.get("reason") == "rewrite_during_generation"
for payload in reset_events
)
assert all(payload["prompt"] != "segment two" for payload in segment_start_events[1:])
reset_events = [payload for payload in ws.sent_json if payload.get("type") == "seed_prompts_reset_applied"]
assert any(payload.get("reason") == "rewrite_during_generation" for payload in reset_events)
finally:
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
mock_server.LATENCY_MS = old_latency_ms
@@ -247,24 +219,24 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 1
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "custom_editable",
"preset_label": "Custom rollout",
"curated_prompts": [],
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.20, {"type": "leave"}),
]
)
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "custom_editable",
"preset_label": "Custom rollout",
"curated_prompts": [],
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.20, {
"type": "leave"
}),
])
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -276,15 +248,9 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
assert "ltx2_stream_start" in message_types
assert "prompt_sources_blocked" not in message_types
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
assert segment_start_events
assert segment_start_events[0]["prompt"] == (
"A moonbase corridor thriller with flooding [segment 1]"
)
assert segment_start_events[0]["prompt"] == ("A moonbase corridor thriller with flooding [segment 1]")
finally:
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
mock_server.LATENCY_MS = old_latency_ms
@@ -297,35 +263,37 @@ def test_mock_server_can_start_new_project_without_reconnecting():
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 40
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.02, {"type": "end_project_keep_session"}),
(
0.20,
{
"type": "project_init_v1",
"preset_id": "test_preset_2",
"preset_label": "Test Preset 2",
"curated_prompts": ["segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.40, {"type": "leave"}),
]
)
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.02, {
"type": "end_project_keep_session"
}),
(
0.20,
{
"type": "project_init_v1",
"preset_id": "test_preset_2",
"preset_label": "Test Preset 2",
"curated_prompts": ["segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.40, {
"type": "leave"
}),
])
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -336,16 +304,11 @@ def test_mock_server_can_start_new_project_without_reconnecting():
project_idle_index = message_types.index("project_idle")
stream_start_indexes = [
index for index, message_type in enumerate(message_types)
if message_type == "ltx2_stream_start"
index for index, message_type in enumerate(message_types) if message_type == "ltx2_stream_start"
]
assert stream_start_indexes[0] < project_idle_index < stream_start_indexes[1]
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
"segment one",
"segment two",
@@ -6,7 +6,6 @@ import os
import re
import time
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -22,6 +21,7 @@ from dreamverse.prompt_enhancer import (
class _FakeResponse:
def __init__(self, payload: dict):
self._payload = payload
@@ -30,6 +30,7 @@ class _FakeResponse:
class _FakeSyncCompletions:
def __init__(self, payload: dict):
self._payload = payload
@@ -38,6 +39,7 @@ class _FakeSyncCompletions:
class _FakeSyncClient:
def __init__(self, payload: dict):
self.chat = type(
"_FakeChat",
@@ -47,6 +49,7 @@ class _FakeSyncClient:
class _DelayedSyncCompletions:
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
self._payload = payload
self._delay_s = delay_s
@@ -61,29 +64,26 @@ class _DelayedSyncCompletions:
class _DelayedSyncClient:
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
self.chat = type(
"_FakeChat",
(),
{
"completions": _DelayedSyncCompletions(
payload,
delay_s=delay_s,
exc=exc,
)
},
{"completions": _DelayedSyncCompletions(
payload,
delay_s=delay_s,
exc=exc,
)},
)()
def _chat_payload_with_content(content: str) -> dict:
return {
"choices": [
{
"message": {
"content": content,
}
"choices": [{
"message": {
"content": content,
}
]
}]
}
@@ -172,6 +172,7 @@ def _build_staged_enhancer(
class _FakeOpenAIClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.chat = type(
@@ -182,6 +183,7 @@ class _FakeOpenAIClient:
class _FakeCerebrasClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.chat = type(
@@ -192,16 +194,12 @@ class _FakeCerebrasClient:
def test_parse_json_response_accepts_fenced_json_with_prose():
parsed = _parse_json_response(
"Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks."
)
parsed = _parse_json_response("Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks.")
assert parsed == {"segment_prompts": ["A", "B"]}
def test_parse_json_response_extracts_first_embedded_object():
parsed = _parse_json_response(
"Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)"
)
parsed = _parse_json_response("Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)")
assert parsed == {"segment_prompts": ["A", "B"]}
@@ -268,16 +266,12 @@ def test_build_client_supports_groq_provider(monkeypatch):
def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
)
)
))
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "preset_a"
@@ -286,15 +280,12 @@ def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"rewritten_prompts":["A","B"]}')
)
enhancer = _build_test_enhancer(_chat_payload_with_content('{"rewritten_prompts":["A","B"]}'))
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
)
)
))
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "current_rollout"
@@ -303,19 +294,14 @@ def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
def test_rewrite_prompt_sequence_accepts_segment_dicts_without_top_level_rollout_metadata():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"segments":[{"prompt":"A"},{"text":"B"}]}'
)
)
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segments":[{"prompt":"A"},{"text":"B"}]}'))
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
preset_id="preset_a",
preset_label="Preset A",
rewrite_instruction="make it cinematic",
)
)
))
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "preset_a"
@@ -329,14 +315,12 @@ def test_rewrite_prompt_sequence_accepts_numbered_prose_output():
"The user is asking for a cinematic rewrite.\n\n"
"1. A dog bounds across the moon's dusty surface, kicking up silver regolith as it chases a rabbit beneath the black sky.\n"
"2. The rabbit darts around a crater rim while the dog lunges after it, Earth glowing blue in the distance.\n"
)
)
))
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
)
)
))
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "current_rollout"
@@ -355,12 +339,10 @@ def test_enhance_prompt_prefers_cerebras_before_groq_fallback():
groq_delay_s=0.01,
)
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
)
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
))
assert result.fallback_used is False
assert result.error is None
@@ -381,12 +363,10 @@ def test_enhance_prompt_uses_groq_when_cerebras_fails():
groq_delay_s=0.01,
)
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
)
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
))
assert result.fallback_used is False
assert result.error is None
@@ -408,13 +388,11 @@ def test_enhance_prompt_can_use_groq_when_cerebras_times_out():
enhancer.http_timeout_ms = 50
enhancer.default_timeout_ms = 50
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
timeout_ms=50,
)
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
timeout_ms=50,
))
assert result.fallback_used is False
assert result.error is None
@@ -434,12 +412,10 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
groq_delay_s=0.08,
)
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
)
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
))
assert result.fallback_used is False
assert result.error is None
@@ -453,15 +429,12 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
enhancer = _build_test_enhancer(
_chat_payload_with_content("I cannot comply with JSON right now.")
)
enhancer = _build_test_enhancer(_chat_payload_with_content("I cannot comply with JSON right now."))
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
)
)
))
assert result.fallback_used is True
assert "No JSON object found in assistant response." in (result.error or "")
assert result.raw_response_text == "I cannot comply with JSON right now."
@@ -473,9 +446,7 @@ def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
)
)
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'))
captured = {
"body": None,
"timeout_seconds": None,
@@ -486,8 +457,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
captured["timeout_seconds"] = timeout_seconds
return (
_chat_payload_with_content(
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
),
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'),
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}',
)
@@ -502,8 +472,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
rewrite_model="gpt-test",
rewrite_temperature=0.2,
timeout_ms=800,
)
)
))
assert result.fallback_used is False
assert captured["body"]["messages"][0] == {
@@ -512,12 +481,12 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
}
assert captured["body"]["messages"][1]["role"] == "user"
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
"mode": "edit_existing_rollout",
"request": (
"Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."
),
"user_instruction": "make it cinematic",
"mode":
"edit_existing_rollout",
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."),
"user_instruction":
"make it cinematic",
"current_rollout": {
"id": "preset_a",
"label": "Preset A",
@@ -528,11 +497,8 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
def test_rewrite_prompt_sequence_supports_new_rollout_mode():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'
)
)
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'))
captured = {
"body": None,
}
@@ -541,10 +507,8 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
del timeout_seconds
captured["body"] = body
return (
_chat_payload_with_content(
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'
),
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'),
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}',
)
@@ -560,30 +524,29 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
rewrite_model="gpt-test",
rewrite_temperature=0.2,
timeout_ms=800,
)
)
))
assert result.fallback_used is False
assert result.prompts == ["A", "B", "C", "D", "E", "F"]
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
"mode": "new_rollout",
"request": (
"Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."
),
"user_instruction": "A moonbase corridor thriller with flooding and red alarms",
"desired_segment_count": 6,
"rollout_id_hint": "custom_editable",
"rollout_label_hint": "Custom rollout",
"mode":
"new_rollout",
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."),
"user_instruction":
"A moonbase corridor thriller with flooding and red alarms",
"desired_segment_count":
6,
"rollout_id_hint":
"custom_editable",
"rollout_label_hint":
"Custom rollout",
}
def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
enhancer.rewrite_all_system_prompt = "shared system prompt"
captured = {
"body": None,
@@ -593,9 +556,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
del timeout_seconds
captured["body"] = body
return (
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
),
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'),
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}',
)
@@ -609,8 +570,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
rewrite_instruction="make it cinematic",
rewrite_model="gpt-test",
system_prompt_override="session specific system prompt",
)
)
))
assert result.fallback_used is False
assert captured["body"]["messages"][0] == {
@@ -621,10 +581,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
@@ -635,24 +592,17 @@ def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
def test_resolve_rewrite_new_rollout_system_prompt_prefers_override():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt(
"session specific system prompt"
)
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt("session specific system prompt")
assert resolved == "session specific system prompt"
def test_generate_auto_prompt_uses_selected_model():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"next_prompt":"Auto next"}')
)
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Auto next"}'))
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
enhancer.rewrite_default_model = "gpt-test"
@@ -680,8 +630,7 @@ def test_generate_auto_prompt_uses_selected_model():
next_segment_idx=2,
model="gpt-alt",
timeout_ms=800,
)
)
))
assert result.fallback_used is False
assert result.error is None
assert result.prompt == "Auto next"
@@ -690,9 +639,7 @@ def test_generate_auto_prompt_uses_selected_model():
def test_enhance_prompt_uses_selected_model():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"next_prompt":"Enhanced next"}')
)
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Enhanced next"}'))
enhancer.enhance_system_prompt = "enhance system prompt"
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
@@ -722,8 +669,7 @@ def test_enhance_prompt_uses_selected_model():
next_segment_idx=2,
model="gpt-alt",
timeout_ms=800,
)
)
))
assert result.fallback_used is False
assert result.error is None
assert result.prompt == "Enhanced next"
@@ -732,9 +678,7 @@ def test_enhance_prompt_uses_selected_model():
def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"prompt":"Extended single clip"}')
)
enhancer = _build_test_enhancer(_chat_payload_with_content('{"prompt":"Extended single clip"}'))
enhancer.enhance_system_prompt = "enhance system prompt"
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
@@ -764,14 +708,12 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
enhancer._request_content = _fake_request_content # type: ignore[attr-defined]
result = asyncio.run(
enhancer.enhance_prompt(
"short 5s idea",
mode="single_clip",
model="gpt-alt",
timeout_ms=800,
)
)
result = asyncio.run(enhancer.enhance_prompt(
"short 5s idea",
mode="single_clip",
model="gpt-alt",
timeout_ms=800,
))
assert result.fallback_used is False
assert result.error is None
assert result.prompt == "Extended single clip"
@@ -784,17 +726,15 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
"single 5-second LTX-2.3 video clip. Respond with "
'valid JSON only as {"prompt": "..."}.' # noqa: E501
),
"user_prompt": "short 5s idea",
"user_prompt":
"short 5s idea",
}
def test_enhance_prompt_single_clip_rejects_plain_text_response():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
"Medium shot of a woman by a rainy cafe window as she lifts her "
"phone, exhales softly, and the camera makes a slow push in."
)
)
_chat_payload_with_content("Medium shot of a woman by a rainy cafe window as she lifts her "
"phone, exhales softly, and the camera makes a slow push in."))
enhancer.auto_system_prompt = "auto system prompt"
result = asyncio.run(
@@ -803,17 +743,14 @@ def test_enhance_prompt_single_clip_rejects_plain_text_response():
mode="single_clip",
model="gpt-test",
timeout_ms=800,
)
)
))
assert result.fallback_used is True
assert "No JSON object found in assistant response." in result.error
assert result.prompt == ""
def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"segment_prompts":["A","B"]}')
)
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segment_prompts":["A","B"]}'))
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test"]
enhancer.rewrite_default_model = "gpt-test"
@@ -824,17 +761,14 @@ def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
mode="single_clip",
model="gpt-test",
timeout_ms=800,
)
)
))
assert result.fallback_used is True
assert result.prompt == ""
assert "Missing prompt string." in (result.error or "")
def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
enhancer = _build_test_enhancer(
_chat_payload_with_content("A cinematic continuation with slow dolly movement.")
)
enhancer = _build_test_enhancer(_chat_payload_with_content("A cinematic continuation with slow dolly movement."))
enhancer.enhance_system_prompt = "enhance system prompt"
enhancer.rewrite_model_options = ["gpt-test"]
enhancer.rewrite_default_model = "gpt-test"
@@ -846,17 +780,14 @@ def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
next_segment_idx=2,
model="gpt-test",
timeout_ms=800,
)
)
))
assert result.fallback_used is True
assert result.prompt == ""
assert "No JSON object found in assistant response." in (result.error or "")
def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
enhancer = _build_test_enhancer(
_chat_payload_with_content("A calm, grounded continuation with subtle motion.")
)
enhancer = _build_test_enhancer(_chat_payload_with_content("A calm, grounded continuation with subtle motion."))
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test"]
enhancer.rewrite_default_model = "gpt-test"
@@ -867,34 +798,30 @@ def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
next_segment_idx=2,
model="gpt-test",
timeout_ms=800,
)
)
))
assert result.fallback_used is True
assert result.prompt == ""
assert "No JSON object found in assistant response." in (result.error or "")
def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
enhancer = _build_test_enhancer(
{
"choices": [
{
"finish_reason": "length",
"message": {
"content": [],
"refusal": None,
},
}
],
"usage": {"completion_tokens": 0},
}
)
enhancer = _build_test_enhancer({
"choices": [{
"finish_reason": "length",
"message": {
"content": [],
"refusal": None,
},
}],
"usage": {
"completion_tokens": 0
},
})
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
)
)
))
assert result.fallback_used is True
assert "No rewrite segment prompts found in assistant response." in (result.error or "")
assert isinstance(result.raw_response_text, str)
@@ -903,10 +830,7 @@ def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
def test_get_rewrite_model_config_returns_fixed_defaults():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
enhancer.rewrite_default_model = "gpt-oss-120b"
enhancer.rewrite_model_options = ["gpt-oss-120b"]
@@ -918,10 +842,7 @@ def test_get_rewrite_model_config_returns_fixed_defaults():
def test_get_prompt_config_includes_auto_extension_prompt():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
enhancer.enhance_system_prompt_path = "/tmp/next.md"
enhancer.auto_system_prompt_path = "/tmp/auto.md"
enhancer.rewrite_all_system_prompt_path = "/tmp/rewrite.md"
@@ -948,19 +869,14 @@ def test_get_prompt_config_includes_auto_extension_prompt():
def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
rewrite_fallback_path = tmp_path / "rewrite_window_system_prompt.md"
rewrite_fallback_path.write_text("rewrite prompt\n", encoding="utf-8")
next_path = tmp_path / "next.md"
next_path.write_text("next prompt\n", encoding="utf-8")
auto_path = tmp_path / "auto.md"
auto_path.write_text("auto prompt\n", encoding="utf-8")
enhancer.rewrite_all_system_prompt_path = str(
tmp_path / "prompts.local" / "rewrite_window_system_prompt.md"
)
enhancer.rewrite_all_system_prompt_path = str(tmp_path / "prompts.local" / "rewrite_window_system_prompt.md")
enhancer.rewrite_all_system_prompt_fallback_path = str(rewrite_fallback_path)
enhancer.enhance_system_prompt_path = str(next_path)
enhancer.auto_system_prompt_path = str(auto_path)
@@ -973,14 +889,9 @@ def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
assert config["rewrite_window_system_prompt_path"] == str(rewrite_fallback_path)
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(
tmp_path,
):
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(tmp_path, ):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
@@ -1007,10 +918,7 @@ def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_emp
def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -1021,9 +929,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
enhancer.auto_system_prompt_path = str(auto_path)
enhancer.rewrite_all_system_prompt_path = str(rewrite_path)
config = enhancer.save_prompt_config(
auto_extension_system_prompt="auto updated",
)
config = enhancer.save_prompt_config(auto_extension_system_prompt="auto updated", )
assert auto_path.read_text(encoding="utf-8").strip() == "auto updated"
assert config["auto_extension_system_prompt"] == "auto updated"
@@ -1031,10 +937,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -1052,9 +955,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
enhancer.rewrite_all_system_prompt_fallback_path = None
enhancer.rewrite_user_system_prompt_fallback_path = None
config = enhancer.save_prompt_config(
rewrite_user_system_prompt="rewrite user updated",
)
config = enhancer.save_prompt_config(rewrite_user_system_prompt="rewrite user updated", )
assert rewrite_user_path.read_text(encoding="utf-8").strip() == "rewrite user updated"
assert config["rewrite_user_system_prompt"] == "rewrite user updated"
@@ -1062,10 +963,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
def test_save_prompt_config_updates_rewrite_model(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -1081,9 +979,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
enhancer.rewrite_default_model = "gpt-test"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
config = enhancer.save_prompt_config(
rewrite_model="gpt-alt",
)
config = enhancer.save_prompt_config(rewrite_model="gpt-alt", )
assert enhancer.rewrite_default_model == "gpt-alt"
assert config["rewrite_model"] == "gpt-alt"
@@ -1092,10 +988,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -1109,9 +1002,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
enhancer.auto_system_prompt_fallback_path = None
enhancer.rewrite_all_system_prompt_fallback_path = None
config = enhancer.save_prompt_config(
rewrite_temperature=1.3,
)
config = enhancer.save_prompt_config(rewrite_temperature=1.3, )
assert enhancer.rewrite_default_temperature == 1.3
assert config["rewrite_temperature"] == 1.3
@@ -1119,10 +1010,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
@@ -1136,13 +1024,9 @@ def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_pat
enhancer.auto_system_prompt_fallback_path = None
enhancer.rewrite_all_system_prompt_fallback_path = None
enhancer.save_prompt_config(
rewrite_window_system_prompt="rewrite updated",
)
enhancer.save_prompt_config(rewrite_window_system_prompt="rewrite updated", )
backup_paths = sorted(
tmp_path.glob("rewrite_window_system_prompt.*.bak.md")
)
backup_paths = sorted(tmp_path.glob("rewrite_window_system_prompt.*.bak.md"))
assert rewrite_path.read_text(encoding="utf-8").strip() == "rewrite updated"
assert len(backup_paths) == 1
@@ -27,13 +27,8 @@ try:
except ModuleNotFoundError:
websockets = None # type: ignore[assignment]
DEFAULT_PRESET_FILE = (
Path(__file__).resolve().parents[2]
/ "web"
/ "prompts"
/ "selected_ltx2_continuation_story_presets.json"
)
DEFAULT_PRESET_FILE = (Path(__file__).resolve().parents[2] / "web" / "prompts" /
"selected_ltx2_continuation_story_presets.json")
def utc_now_iso() -> str:
@@ -65,10 +60,7 @@ def safe_percentile(values: list[float], percentile: float) -> float | None:
if lower == upper:
return sorted_values[lower]
fraction = rank - lower
return (
sorted_values[lower]
+ (sorted_values[upper] - sorted_values[lower]) * fraction
)
return (sorted_values[lower] + (sorted_values[upper] - sorted_values[lower]) * fraction)
def summarize_series(values: list[float]) -> dict[str, float | int | None]:
@@ -145,24 +137,16 @@ def load_curated_prompts(
selected_id = str(selected.get("id", "")).strip() or "unknown_preset"
raw_prompts = selected.get("segment_prompts", [])
if not isinstance(raw_prompts, list):
raise ValueError(
f"Preset {selected_id} has invalid segment_prompts (must be list)."
)
raise ValueError(f"Preset {selected_id} has invalid segment_prompts (must be list).")
prompts = [
str(prompt).strip()
for prompt in raw_prompts
if isinstance(prompt, str) and str(prompt).strip()
]
prompts = [str(prompt).strip() for prompt in raw_prompts if isinstance(prompt, str) and str(prompt).strip()]
if not prompts:
raise ValueError(f"Preset {selected_id} has no non-empty prompts.")
limited = prompts[:curated_limit]
if not limited:
raise ValueError(
f"curated_limit={curated_limit} produced no prompts for preset "
f"{selected_id}."
)
raise ValueError(f"curated_limit={curated_limit} produced no prompts for preset "
f"{selected_id}.")
return selected_id, limited, len(prompts)
@@ -224,11 +208,11 @@ async def run_single_session(
try:
async with websockets.connect(
url,
max_size=None,
ping_interval=None,
open_timeout=connect_timeout_s,
close_timeout=2.0,
url,
max_size=None,
ping_interval=None,
open_timeout=connect_timeout_s,
close_timeout=2.0,
) as ws:
connect_finish_monotonic = time.monotonic()
session_data["connect_finish_ts_utc"] = utc_now_iso()
@@ -249,9 +233,7 @@ async def run_single_session(
timeout_remaining = session_timeout_s - elapsed_s
if timeout_remaining <= 0:
session_data["status"] = "timeout"
session_data["error"] = (
f"Session timed out after {session_timeout_s:.1f}s."
)
session_data["error"] = (f"Session timed out after {session_timeout_s:.1f}s.")
break
recv_start_epoch = time.time()
@@ -265,9 +247,7 @@ async def run_single_session(
)
except asyncio.TimeoutError:
session_data["status"] = "timeout"
session_data["error"] = (
"Timed out waiting for websocket message."
)
session_data["error"] = ("Timed out waiting for websocket message.")
break
except Exception as exc:
session_data["status"] = "failed"
@@ -288,20 +268,16 @@ async def run_single_session(
chunk_gap_ms: float | None = None
if last_chunk_finish_monotonic is not None:
chunk_gap_ms = (
recv_finish_monotonic - last_chunk_finish_monotonic
) * 1000.0
chunk_gap_ms = (recv_finish_monotonic - last_chunk_finish_monotonic) * 1000.0
session_data["chunks"].append(
{
"segment_idx": current_segment_idx,
"chunk_idx": session_data["total_chunks"],
"size_bytes": len(message),
"chunk_start_ts_utc": recv_start_iso,
"chunk_finish_ts_utc": recv_finish_iso,
"chunk_gap_ms": chunk_gap_ms,
}
)
session_data["chunks"].append({
"segment_idx": current_segment_idx,
"chunk_idx": session_data["total_chunks"],
"size_bytes": len(message),
"chunk_start_ts_utc": recv_start_iso,
"chunk_finish_ts_utc": recv_finish_iso,
"chunk_gap_ms": chunk_gap_ms,
})
last_chunk_finish_monotonic = recv_finish_monotonic
last_chunk_finish_epoch = recv_finish_epoch
session_data["last_chunk_finish_ts_utc"] = recv_finish_iso
@@ -321,9 +297,7 @@ async def run_single_session(
if msg_type == "gpu_assigned":
session_data["gpu_assigned_ts_utc"] = recv_finish_iso
if connect_finish_monotonic is not None:
session_data["queue_wait_ms"] = (
recv_finish_monotonic - connect_finish_monotonic
) * 1000.0
session_data["queue_wait_ms"] = (recv_finish_monotonic - connect_finish_monotonic) * 1000.0
elif msg_type == "ltx2_stream_start":
if initial_total_segments is None:
parsed_total = parse_int(data.get("total_segments"))
@@ -338,20 +312,13 @@ async def run_single_session(
session_data["media_segments_completed"] += 1
if first_media_segment_complete_epoch is None:
first_media_segment_complete_epoch = recv_finish_epoch
session_data[
"first_media_segment_complete_ts_utc"
] = recv_finish_iso
session_data["first_media_segment_complete_ts_utc"] = recv_finish_iso
elif msg_type == "ltx2_segment_complete":
session_data["segments_completed"] += 1
seg_idx = parse_int(data.get("segment_idx"))
if (
initial_total_segments is not None
and seg_idx is not None
and seg_idx >= initial_total_segments
):
session_data[
"target_segment_complete_ts_utc"
] = recv_finish_iso
if (initial_total_segments is not None and seg_idx is not None
and seg_idx >= initial_total_segments):
session_data["target_segment_complete_ts_utc"] = recv_finish_iso
await asyncio.sleep(post_complete_wait_s)
session_data["leave_sent_ts_utc"] = utc_now_iso()
try:
@@ -362,15 +329,11 @@ async def run_single_session(
break
elif msg_type == "session_timeout":
session_data["status"] = "timeout"
session_data["error"] = str(
data.get("message") or "Backend session timeout"
)
session_data["error"] = str(data.get("message") or "Backend session timeout")
break
elif msg_type == "error":
session_data["status"] = "failed"
session_data["error"] = str(
data.get("message") or "Backend error message"
)
session_data["error"] = str(data.get("message") or "Backend error message")
break
if session_data["status"] == "failed" and session_data["error"] is None:
@@ -379,29 +342,18 @@ async def run_single_session(
session_data["status"] = "failed"
session_data["error"] = f"WebSocket connect/run failed: {exc}"
if (
first_chunk_finish_epoch is not None
and last_chunk_finish_epoch is not None
and session_data["total_chunk_bytes"] > 0
):
if (first_chunk_finish_epoch is not None and last_chunk_finish_epoch is not None
and session_data["total_chunk_bytes"] > 0):
duration_s = last_chunk_finish_epoch - first_chunk_finish_epoch
if duration_s > 0:
session_data["session_goodput_mbps"] = (
session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0
)
session_data["session_goodput_mbps"] = (session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0)
if (
first_chunk_finish_epoch is not None
and first_media_segment_complete_epoch is not None
):
session_data["first_chunk_before_first_media_complete"] = (
first_chunk_finish_epoch < first_media_segment_complete_epoch
)
if (first_chunk_finish_epoch is not None and first_media_segment_complete_epoch is not None):
session_data["first_chunk_before_first_media_complete"] = (first_chunk_finish_epoch
< first_media_segment_complete_epoch)
session_data["close_ts_utc"] = utc_now_iso()
session_data["duration_ms"] = (
time.monotonic() - session_start_monotonic
) * 1000.0
session_data["duration_ms"] = (time.monotonic() - session_start_monotonic) * 1000.0
return session_data
@@ -412,14 +364,11 @@ async def run_worker_sessions(
config: dict[str, Any],
) -> list[dict[str, Any]]:
tasks = [
asyncio.create_task(
run_single_session(
worker_id=worker_id,
worker_session_idx=idx,
config=config,
)
)
for idx in range(session_count)
asyncio.create_task(run_single_session(
worker_id=worker_id,
worker_session_idx=idx,
config=config,
)) for idx in range(session_count)
]
if not tasks:
return []
@@ -437,29 +386,23 @@ def worker_entry(
try:
ready_queue.put({"worker_id": worker_id, "status": "ready"})
start_event.wait()
sessions = asyncio.run(
run_worker_sessions(
worker_id=worker_id,
session_count=session_count,
config=config,
)
)
result_queue.put(
{
"worker_id": worker_id,
"status": "ok",
"sessions": sessions,
}
)
sessions = asyncio.run(run_worker_sessions(
worker_id=worker_id,
session_count=session_count,
config=config,
))
result_queue.put({
"worker_id": worker_id,
"status": "ok",
"sessions": sessions,
})
except Exception as exc:
result_queue.put(
{
"worker_id": worker_id,
"status": "error",
"error": str(exc),
"traceback": traceback.format_exc(),
}
)
result_queue.put({
"worker_id": worker_id,
"status": "error",
"error": str(exc),
"traceback": traceback.format_exc(),
})
def build_summary(
@@ -517,33 +460,22 @@ def build_summary(
if len(all_chunk_finish_epochs) >= 2 and total_chunk_bytes > 0:
duration_s = max(all_chunk_finish_epochs) - min(all_chunk_finish_epochs)
if duration_s > 0:
global_goodput_mbps = (
total_chunk_bytes * 8.0 / duration_s / 1_000_000.0
)
global_goodput_mbps = (total_chunk_bytes * 8.0 / duration_s / 1_000_000.0)
bucket_throughputs_mbps = [
(bytes_count * 8.0) / 1_000_000.0
for _, bytes_count in sorted(bucket_bytes.items())
]
bucket_throughputs_mbps = [(bytes_count * 8.0) / 1_000_000.0 for _, bytes_count in sorted(bucket_bytes.items())]
bucket_stats = summarize_series(bucket_throughputs_mbps)
chunk_gap_threshold_breaches = [
value for value in chunk_gaps if value >= chunk_gap_threshold_ms
]
chunk_gap_threshold_breaches = [value for value in chunk_gaps if value >= chunk_gap_threshold_ms]
non_success = len(sessions) - status_counts.get("success", 0)
fail_reasons: list[str] = []
if non_success > 0:
fail_reasons.append(
f"{non_success} session(s) did not complete successfully."
)
fail_reasons.append(f"{non_success} session(s) did not complete successfully.")
if not chunk_gaps:
fail_reasons.append("No chunk gap data collected.")
if chunk_gap_threshold_breaches:
fail_reasons.append(
f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
f"{chunk_gap_threshold_ms:.0f}ms."
)
fail_reasons.append(f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
f"{chunk_gap_threshold_ms:.0f}ms.")
passed = len(fail_reasons) == 0
progressive_ratio = None
@@ -554,20 +486,18 @@ def build_summary(
"passed": passed,
"fail_reasons": fail_reasons,
"sessions": {
"total": len(sessions),
"success": status_counts.get("success", 0),
"failed": status_counts.get("failed", 0),
"timeout": status_counts.get("timeout", 0),
"protocol_error": status_counts.get("protocol_error", 0),
"other": (
len(sessions)
- (
status_counts.get("success", 0)
+ status_counts.get("failed", 0)
+ status_counts.get("timeout", 0)
+ status_counts.get("protocol_error", 0)
)
),
"total":
len(sessions),
"success":
status_counts.get("success", 0),
"failed":
status_counts.get("failed", 0),
"timeout":
status_counts.get("timeout", 0),
"protocol_error":
status_counts.get("protocol_error", 0),
"other": (len(sessions) - (status_counts.get("success", 0) + status_counts.get("failed", 0) +
status_counts.get("timeout", 0) + status_counts.get("protocol_error", 0))),
},
"chunk_gap_ms": {
**chunk_gap_stats,
@@ -606,51 +536,39 @@ def print_summary(
bucket_bw = bandwidth["bucketed_1s"]
print("=== LTX2 Realtime Stress Test Summary ===")
print(
"Run: "
f"url={run_info['url']} clients={run_info['clients']} "
f"processes={run_info['processes']} "
f"preset={run_info['preset_id']} "
f"curated_limit={run_info['curated_limit']}"
)
print(
"Sessions: "
f"total={sessions['total']} success={sessions['success']} "
f"failed={sessions['failed']} timeout={sessions['timeout']} "
f"protocol_error={sessions['protocol_error']}"
)
print(
"Chunk gap ms: "
f"min={format_num(chunk_gap['min'])} "
f"p50={format_num(chunk_gap['p50'])} "
f"p95={format_num(chunk_gap['p95'])} "
f"p99={format_num(chunk_gap['p99'])} "
f"max={format_num(chunk_gap['max'])} "
f"threshold={format_num(chunk_gap['threshold_ms'])} "
f"breaches={chunk_gap['breach_count']}"
)
print(
"Queue wait ms: "
f"min={format_num(queue_wait['min'])} "
f"p50={format_num(queue_wait['p50'])} "
f"p95={format_num(queue_wait['p95'])} "
f"max={format_num(queue_wait['max'])}"
)
print("Run: "
f"url={run_info['url']} clients={run_info['clients']} "
f"processes={run_info['processes']} "
f"preset={run_info['preset_id']} "
f"curated_limit={run_info['curated_limit']}")
print("Sessions: "
f"total={sessions['total']} success={sessions['success']} "
f"failed={sessions['failed']} timeout={sessions['timeout']} "
f"protocol_error={sessions['protocol_error']}")
print("Chunk gap ms: "
f"min={format_num(chunk_gap['min'])} "
f"p50={format_num(chunk_gap['p50'])} "
f"p95={format_num(chunk_gap['p95'])} "
f"p99={format_num(chunk_gap['p99'])} "
f"max={format_num(chunk_gap['max'])} "
f"threshold={format_num(chunk_gap['threshold_ms'])} "
f"breaches={chunk_gap['breach_count']}")
print("Queue wait ms: "
f"min={format_num(queue_wait['min'])} "
f"p50={format_num(queue_wait['p50'])} "
f"p95={format_num(queue_wait['p95'])} "
f"max={format_num(queue_wait['max'])}")
ratio = progressive["ratio"]
ratio_text = "n/a" if ratio is None else f"{ratio * 100:.2f}%"
print(
"Progressive streaming: "
f"{progressive['success_sessions']}/"
f"{progressive['eligible_sessions']} ({ratio_text})"
)
print(
"Bandwidth Mbps: "
f"per_session_avg={format_num(per_session_bw['avg'])} "
f"per_session_p95={format_num(per_session_bw['p95'])} "
f"global={format_num(bandwidth['global_goodput_mbps'])} "
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}"
)
print("Progressive streaming: "
f"{progressive['success_sessions']}/"
f"{progressive['eligible_sessions']} ({ratio_text})")
print("Bandwidth Mbps: "
f"per_session_avg={format_num(per_session_bw['avg'])} "
f"per_session_p95={format_num(per_session_bw['p95'])} "
f"global={format_num(bandwidth['global_goodput_mbps'])} "
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}")
print(f"VERDICT: {'PASS' if summary['passed'] else 'FAIL'}")
if summary["fail_reasons"]:
print("Fail reasons:")
@@ -670,10 +588,8 @@ def distribute_sessions(total_clients: int, process_count: int) -> list[int]:
def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
if websockets is None:
raise RuntimeError(
"Missing dependency: websockets. Install it before running this "
"stress test."
)
raise RuntimeError("Missing dependency: websockets. Install it before running this "
"stress test.")
preset_file = Path(args.preset_file).expanduser().resolve()
selected_preset_id, curated_prompts, total_prompt_count = load_curated_prompts(
@@ -735,13 +651,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
start_event.set()
result_deadline = (
time.monotonic()
+ args.connect_timeout_s
+ args.session_timeout_s
+ args.post_complete_wait_s
+ 180.0
)
result_deadline = (time.monotonic() + args.connect_timeout_s + args.session_timeout_s +
args.post_complete_wait_s + 180.0)
worker_results: list[dict[str, Any]] = []
while len(worker_results) < len(processes):
timeout_s = max(0.1, result_deadline - time.monotonic())
@@ -765,24 +676,20 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
if result.get("status") == "ok":
sessions.extend(result.get("sessions", []))
else:
worker_errors.append(
{
"worker_id": result.get("worker_id"),
"error": result.get("error"),
"traceback": result.get("traceback"),
}
)
worker_errors.append({
"worker_id": result.get("worker_id"),
"error": result.get("error"),
"traceback": result.get("traceback"),
})
received_workers = {result.get("worker_id") for result in worker_results}
expected_workers = set(range(len(processes)))
missing_workers = sorted(expected_workers - received_workers)
for worker_id in missing_workers:
worker_errors.append(
{
"worker_id": worker_id,
"error": "No worker result received.",
}
)
worker_errors.append({
"worker_id": worker_id,
"error": "No worker result received.",
})
run_end_epoch = time.time()
run_end_iso = iso_from_epoch(run_end_epoch)
@@ -795,9 +702,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
if worker_errors:
summary["passed"] = False
summary["fail_reasons"] = list(summary["fail_reasons"]) + [
f"{len(worker_errors)} worker error(s) occurred."
]
summary["fail_reasons"] = list(
summary["fail_reasons"]) + [f"{len(worker_errors)} worker error(s) occurred."]
output_payload = {
"run_info": {
@@ -833,9 +739,7 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Multiprocess realtime stress test for LTX2 streaming.",
)
parser = argparse.ArgumentParser(description="Multiprocess realtime stress test for LTX2 streaming.", )
parser.add_argument(
"-u",
"--url",
@@ -47,13 +47,11 @@ def test_persist_session_init_image_returns_none_when_missing_data():
def test_persist_session_init_image_rejects_unsupported_mime():
with pytest.raises(ValueError, match="PNG, JPEG, or WebP"):
persist_session_init_image(
{
"name": "frame.gif",
"mime_type": "image/gif",
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
}
)
persist_session_init_image({
"name": "frame.gif",
"mime_type": "image/gif",
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
})
def test_persist_session_init_image_rejects_large_payload(monkeypatch):
@@ -66,10 +64,8 @@ def test_persist_session_init_image_rejects_large_payload(monkeypatch):
monkeypatch.setattr(base64, "b64decode", fake_b64decode)
with pytest.raises(ValueError, match="15 MB or smaller"):
persist_session_init_image(
{
"name": "frame.png",
"mime_type": "image/png",
"data_url": data_url,
}
)
persist_session_init_image({
"name": "frame.png",
"mime_type": "image/png",
"data_url": data_url,
})
File diff suppressed because it is too large Load Diff
+9 -15
View File
@@ -8,12 +8,10 @@ import modal
IMAGE = os.environ.get("DREAMVERSE_IMAGE")
if not IMAGE:
raise RuntimeError(
"DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag."
)
raise RuntimeError("DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag.")
# ``@modal.web_server`` invokes ``serve()`` directly and bypasses the image
# ENTRYPOINT (``docker/docker_entrypoint.sh``). That entrypoint normally
@@ -65,14 +63,10 @@ def serve():
# ``or ""`` collapses ``None`` (unset) into an empty string, ``.strip()``
# collapses whitespace-only values (e.g. ``" "``) — both should be
# treated as missing.
missing = [
k for k in _REQUIRED_SECRET_KEYS
if not (os.environ.get(k) or "").strip()
]
missing = [k for k in _REQUIRED_SECRET_KEYS if not (os.environ.get(k) or "").strip()]
if missing:
raise RuntimeError(
"dreamverse-api-keys secret is missing required entries: "
f"{', '.join(missing)}. Add them with `modal secret create "
"dreamverse-api-keys ... --force` and redeploy "
"(see apps/dreamverse/scripts/modal/README.md).")
raise RuntimeError("dreamverse-api-keys secret is missing required entries: "
f"{', '.join(missing)}. Add them with `modal secret create "
"dreamverse-api-keys ... --force` and redeploy "
"(see apps/dreamverse/scripts/modal/README.md).")
subprocess.Popen(["dreamverse-server", "--host", "0.0.0.0", "--port", "8009"])
+3 -6
View File
@@ -74,10 +74,8 @@ def test_snapshot_shapes_devices(monkeypatch: pytest.MonkeyPatch) -> None:
}
def test_snapshot_tolerates_missing_sensors(
monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setitem(sys.modules, "pynvml",
_make_fake_pynvml(broken_sensors=True))
def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setitem(sys.modules, "pynvml", _make_fake_pynvml(broken_sensors=True))
snap = gpu_mod.get_gpu_snapshot()
assert snap["available"] is True
g = snap["gpus"][0]
@@ -86,8 +84,7 @@ def test_snapshot_tolerates_missing_sensors(
assert g["power_limit_watts"] is None
def test_snapshot_reports_nvml_failure(
monkeypatch: pytest.MonkeyPatch) -> None:
def test_snapshot_reports_nvml_failure(monkeypatch: pytest.MonkeyPatch) -> None:
fake = _make_fake_pynvml()
fake.nvmlInit = lambda: (_ for _ in ()).throw(_NVMLError("driver gone"))
monkeypatch.setitem(sys.modules, "pynvml", fake)
@@ -130,8 +130,7 @@ def test_dmd_builds_three_role_models_and_method_knobs() -> None:
def test_dmd_vsa_maps_to_training_vsa_sparsity() -> None:
config = build_training_config(
_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
config = build_training_config(_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
assert config["training"]["vsa"]["sparsity"] == 0.9
@@ -161,31 +160,27 @@ def test_validation_callback_only_when_file_given() -> None:
without = build_training_config(_job("full_t2v"), "out")
assert "validation" not in without["callbacks"]
with_file = build_training_config(
_job("full_t2v", validation_dataset_file="val.json"), "out")
with_file = build_training_config(_job("full_t2v", validation_dataset_file="val.json"), "out")
validation = with_file["callbacks"]["validation"]
assert validation["dataset_file"] == "val.json"
assert validation["pipeline_target"].endswith(".WanPipeline")
assert validation["sampling_steps"] == [50]
dmd = build_training_config(
_job("dmd_t2v", validation_dataset_file="val.json"), "out")
dmd = build_training_config(_job("dmd_t2v", validation_dataset_file="val.json"), "out")
validation = dmd["callbacks"]["validation"]
assert validation["pipeline_target"].endswith(".WanDMDPipeline")
assert validation["sampling_steps"] == [3]
assert validation["sampling_timesteps"] == [1000, 757, 522]
# KD/ODE-init has no sampling-based validation pipeline.
ode = build_training_config(
_job("ode_init", validation_dataset_file="val.json"), "out")
ode = build_training_config(_job("ode_init", validation_dataset_file="val.json"), "out")
assert "validation" not in ode["callbacks"]
def test_ltx2_models_are_rejected() -> None:
assert is_ltx2_model("Lightricks/LTX-2-19B")
with pytest.raises(ValueError, match="LTX-2 training is not supported"):
build_training_config(
_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
build_training_config(_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
def test_unknown_workload_is_rejected() -> None:
@@ -195,8 +190,7 @@ def test_unknown_workload_is_rejected() -> None:
def test_invalid_denoising_steps_are_rejected() -> None:
with pytest.raises(ValueError, match="Invalid DMD denoising steps"):
build_training_config(
_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
build_training_config(_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
def test_training_env_has_no_backend_override() -> None:
@@ -211,8 +205,7 @@ def test_workloads_match_frontend_job_config() -> None:
(src/lib/jobConfig.ts) — drift means creatable-but-unrunnable jobs."""
import re
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" /
"jobConfig.ts").read_text(encoding="utf-8")
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" / "jobConfig.ts").read_text(encoding="utf-8")
all_types = set(re.findall(r'type:\s*"([^"]+)"', job_config))
inference_types = {"t2v", "i2v", "t2i"}
assert inference_types <= all_types, "jobConfig.ts parse failed"
+1 -1
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
+52
View File
@@ -0,0 +1,52 @@
{
"recipes": [
{
"id": "fastwan21-t2v",
"task": "Text to video",
"label": "FastWan2.1 1.3B (distilled + VSA)",
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
},
{
"id": "wan22-t2v",
"task": "Text to video",
"label": "Wan2.2 A14B",
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2.py",
"command": "python examples/inference/basic/basic_wan2_2.py"
},
{
"id": "wan21-i2v",
"task": "Image to video",
"label": "Wan2.1 14B 480P",
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
"source": "scripts/inference/inference_wan_i2v.yaml",
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
},
{
"id": "turbowan22-i2v",
"task": "Image to video",
"label": "TurboWan2.2 A14B",
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
},
{
"id": "wan22-ti2v",
"task": "Text or image to video",
"label": "Wan2.2 TI2V 5B",
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
},
{
"id": "matrix-game-2",
"task": "Interactive world",
"label": "Matrix Game 2.0",
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"source": "examples/inference/basic/basic_matrixgame2.py",
"command": "python examples/inference/basic/basic_matrixgame2.py"
}
]
}
+60
View File
@@ -0,0 +1,60 @@
(() => {
let recipesPromise;
const loadRecipes = (url) => {
recipesPromise ||= fetch(url).then((response) => {
if (!response.ok) throw new Error(`HTTP ${response.status}`);
return response.json();
});
return recipesPromise;
};
const init = () => {
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
if (root.dataset.initialized) return;
root.dataset.initialized = "true";
const select = root.querySelector("[data-cookbook-recipe]");
const model = root.querySelector("[data-cookbook-model]");
const source = root.querySelector("[data-cookbook-source]");
const command = root.querySelector("[data-cookbook-command]");
const status = root.querySelector("[data-cookbook-status]");
try {
const { recipes } = await loadRecipes(root.dataset.recipes);
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
const groups = new Map();
select.replaceChildren();
recipes.forEach((recipe) => {
if (!groups.has(recipe.task)) {
const group = document.createElement("optgroup");
group.label = recipe.task;
groups.set(recipe.task, group);
select.append(group);
}
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
});
const render = () => {
const recipe = byId.get(select.value);
model.textContent = recipe.model;
source.textContent = recipe.source;
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
command.textContent = recipe.command;
status.textContent = `${recipe.label} selected.`;
};
select.addEventListener("change", render);
select.disabled = false;
render();
} catch (error) {
status.textContent = "Recipes could not be loaded. Use the examples link below.";
console.error("Failed to load FastVideo cookbook recipes", error);
}
});
};
if (window.document$) window.document$.subscribe(init);
else document.addEventListener("DOMContentLoaded", init);
})();
+40
View File
@@ -42,6 +42,46 @@ img {
margin: 0 auto;
}
.cookbook-picker {
padding: 1rem;
border: 0.05rem solid var(--md-default-fg-color--lightest);
border-radius: 0.2rem;
}
.cookbook-picker select {
width: 100%;
padding: 0.6rem;
color: var(--md-default-fg-color);
background: var(--md-default-bg-color);
border: 0.05rem solid var(--md-default-fg-color--lighter);
border-radius: 0.2rem;
}
.cookbook-picker dl {
display: grid;
grid-template-columns: max-content 1fr;
gap: 0.25rem 1rem;
}
.cookbook-picker dt {
font-weight: 700;
}
.cookbook-picker dd {
margin: 0;
min-width: 0;
overflow-wrap: anywhere;
}
.cookbook-picker__status {
position: absolute;
width: 1px;
height: 1px;
overflow: hidden;
clip: rect(0, 0, 0, 0);
white-space: nowrap;
}
.md-typeset .copy-page-button.md-button {
float: right;
margin: 0 0 1rem 1rem;
+44
View File
@@ -0,0 +1,44 @@
# Inference Cookbook
Choose a complete recipe maintained in the FastVideo repository. Each command
runs its checked-in source directly, so coupled model, GPU, offload, and
attention settings do not drift into unsupported combinations.
The commands expect a local clone:
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
```
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
<label for="cookbook-recipe"><strong>Recipe</strong></label>
<select id="cookbook-recipe" data-cookbook-recipe disabled>
<option>Loading recipes…</option>
</select>
<dl>
<dt>Model</dt>
<dd data-cookbook-model>Loading…</dd>
<dt>Source</dt>
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
</dl>
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
<noscript>
JavaScript is needed for the recipe picker. Browse the
<a href="../inference/examples/examples_inference_index/">inference examples</a>
instead.
</noscript>
</div>
## Customize a recipe
Start from the checked-in source, then change only the settings your model
supports:
- [Configuration](../inference/configuration.md) covers the Python and CLI
config surfaces.
- [Optimizations](../inference/optimizations.md) covers attention backends,
compilation, and memory tradeoffs.
- [Support matrix](../inference/support_matrix.md) lists supported models and
optimizations.
+128
View File
@@ -0,0 +1,128 @@
# Fast mode (RIFE) — Apple Silicon
`--fast` makes local generation ~2.7× faster by **generating fewer frames and
interpolating the rest** with an Apple-Silicon-native RIFE model, instead of
denoising every frame. Video-diffusion denoise is dominated by self-attention,
which is O(tokens²); halving the frames cuts the token count ~2× and the denoise
compute ~3.7×, so the wall-clock drops far more than 2×. RIFE (which estimates
its own optical flow — no motion vectors needed) fills the dropped frames back
in for ~1.4 s, and a light unsharp pass counters its softening.
Measured on the 1.3B INT8 QAD model (fox, 480×832×81, M4): generate 41 + RIFE→81
runs in ~35 s of denoise vs ~90 s full, at reconstruction MS-SSIM **0.97**.
Reproduce with `python -m fastvideo.benchmarks.eval_metalfx_rife --mode int8`.
> **Note:** Apple's *MetalFX* frame interpolation is **not** usable here — it
> requires game-engine motion vectors + depth, which diffusion output lacks. We
> use the video-native **`rife-mlx`** model instead (Metal-backed, torch-free).
## Install
```bash
uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
```
## Use
```bash
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--mlx-checkpoint <FastWan2.1-T2V-1.3B-INT8-QAD> \
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
--num-frames 81 --fast \
--output-path video_samples/fox_fast.mp4
```
`--num-frames` stays the *target* length; fast mode generates the smallest
VAE-aligned keyframe count that RIFE can interpolate to that target.
| Flag | Default | Meaning |
|---|---|---|
| `--fast` / `--no-fast` | off | enable fast mode |
| `--fast-factor` | 2 | generate 1/factor of the frames (2 = half) |
| `--fast-sharpen` | 0.6 | light unsharp strength to counter RIFE softness (0 disables) |
Fast mode composes with everything else (`--mlx-quantization int8`,
`--mlx-compile`, TAEHV vs `--decode-backend wan-vae`). Keep `--fast-factor` at 2
for quality — larger temporal gaps are where RIFE starts inventing motion.
## Spatial fast mode (`--fast-spatial`)
The spatial twin of `--fast`: instead of dropping frames, drop pixels. Denoise
*and decode* at `height/width // fast-spatial-scale`, then resample the decoded
frames up to the requested size. Self-attention is O(tokens²), so halving each
spatial axis cuts the token count 4× and the denoise time far more than that —
measured on the 1.3B INT8 QAD model at 480×832×81, M4 Max: **86.1 s → 10.3 s**
of denoise. It composes with `--fast`; both together run the same clip in
**4.5 s** of denoise.
```bash
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
--height 480 --width 832 --num-frames 81 --fast-spatial \
--output-path video_samples/fox_fast_spatial.mp4
```
| Flag | Default | Meaning |
|---|---|---|
| `--fast-spatial` / `--no-fast-spatial` | off | enable spatial fast mode |
| `--fast-spatial-scale` | 2 | denoise at 1/scale of each spatial axis |
| `--fast-spatial-upsample-mode` | `lanczos` | pixel interpolation kernel (`lanczos`, `cubic`, `bilinear`, `nearest`) |
| `--fast-spatial-sharpen` | 0.4 | light unsharp strength to counter resampling softness (0 disables) |
### The upsample must happen in pixel space
This is the one thing to get right. The obvious implementation — bilinearly
upsample the finished latents and decode at the target size — **does not work**,
and produces a distinctive failure: correct composition and silhouette under a
smeared, hazy veil, with ringing along strong edges.
A Wan latent cell is a *learned code* for an 8×8 (Wan2.1) or 16×16 (Wan2.2)
pixel block, not a low-pass sample of the image. The average of two adjacent
codes is not the code of the averaged blocks; it is a vector the decoder was
never trained on. Measured on Wan2.1-1.3B at 480×832, a 2× bilinear latent
upsample destroys **62%** of the latent's high-frequency energy while leaving
its overall magnitude intact — exactly the signature of that veil. At Wan2.2-5B
the same operation degrades to black or noise.
Decoded RGB frames have no such problem: an image *is* a sampled 2-D signal, so
Lanczos interpolation is the operation it was defined for. The result is soft —
it carries stage-1's real detail budget and no more — but clean and coherent.
`--refine` gets away with a latent-space upsample only because a second DMD pass
re-denoises the hand-off; spatial fast mode passes the latent straight to the
decoder, so it cannot.
## Refine (`--refine`) stage-2 timesteps
`--refine` hands stage 1 to stage 2 as `(1 - sigma) * upsampled + sigma * noise`,
where `sigma` comes from the *first* stage-2 timestep. FastWan's DMD grid opens
at `t=1000`, which is `sigma == 1` exactly — so a stage-2 grid that starts there
weights the stage-1 result at zero and refine silently degrades into a plain
full-resolution run at twice the cost.
Left unset, `--refine-dmd-denoising-steps` now derives the stage-2 grid from the
stage-1 one with leading full-noise steps dropped (`1000,757,522` → `757,522`).
That keeps the pass on timesteps the distilled student was trained on while
letting stage-1 structure through: hand-off `sigma = 0.757`, stage-1 weight
`0.243`. Passing a grid that starts at full noise is now an error rather than a
silently wasted pass.
The run prints the resolved hand-off so it is visible:
```
[refine] stage-2 hand-off sigma=0.7568 (stage-1 weight 0.2432)
```
There is a trade-off in choosing that grid. Later start = more of the draft
survives, but fewer stage-2 steps. On Wan2.1 the default `757,522` gives weight
0.243 with two steps; `--refine-dmd-denoising-steps 522` gives weight 0.478 with
one. `--refine-sigma` decouples the noise level from the timestep entirely — it
logs a warning, because the DiT is then told a timestep that does not match the
noise it receives.
**Wan2.2-5B has a lower ceiling.** Its warped schedule maps `1000,757,522` to
sigmas `1.000, 0.940, 0.845`, so the best available stage-1 weight is **0.060**
(vs 0.243 at 1.3B). Un-warped (`--no-warp`) the same grid gives `1.000, 0.757,
0.522` and a weight of 0.243 — but warping is what matches the FastVideo
sampling schedule, so turning it off changes the timesteps the distilled student
sees. Which is better at 5B is unresolved and needs a run on real 5B weights.
@@ -76,6 +76,10 @@ surfaces:
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
+37
View File
@@ -2,6 +2,7 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import json
import os
import re
from dataclasses import dataclass, field
@@ -19,6 +20,40 @@ GENERATED_DOC_PREFIXES = (
"training/examples/",
"distillation/examples/",
)
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
COOKBOOK_SOURCE_ROOTS = (
ROOT_DIR / "examples/inference",
ROOT_DIR / "scripts/inference",
)
def validate_cookbook() -> None:
"""Keep cookbook entries tied to checked-in runnable sources."""
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
if not isinstance(recipes, list) or not recipes:
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
seen: set[str] = set()
for recipe in recipes:
required = ("id", "task", "label", "model", "source", "command")
missing = {key for key in required if not recipe.get(key)}
if missing:
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
if recipe["id"] in seen:
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
seen.add(recipe["id"])
source = (ROOT_DIR / recipe["source"]).resolve()
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
if not source.is_file():
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
source_text = source.read_text(encoding="utf-8")
if recipe["model"] not in source_text:
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
if recipe["source"] not in recipe["command"]:
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
def fix_case(text: str) -> str:
@@ -536,6 +571,7 @@ def on_pre_build(config, **kwargs):
MkDocs hook to generate examples before building the documentation.
This function is called automatically by MkDocs' native hook system.
"""
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
@@ -549,6 +585,7 @@ def on_page_context(context, page, **kwargs):
if __name__ == "__main__":
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
+3 -2
View File
@@ -65,6 +65,7 @@ uv pip install flash-attn --no-build-isolation -v
## Next Steps
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Quick Start](quick_start.md) - Generate your first video
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
+6 -4
View File
@@ -49,10 +49,12 @@ brew install ffmpeg
### Installation
FastWan's native Apple Silicon runtime requires the `mlx` extra.
#### With uv (recommended)
```bash
uv pip install fastvideo
uv pip install "fastvideo[mlx]"
```
#### With Conda environment (alternative)
@@ -60,7 +62,7 @@ uv pip install fastvideo
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
```bash
uv pip install fastvideo
uv pip install "fastvideo[mlx]"
```
### Installation from Source
@@ -76,13 +78,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Basic installation:
```bash
uv pip install -e .
uv pip install -e ".[mlx]"
```
Alternative with Conda environment:
```bash
uv pip install -e .
uv pip install -e ".[mlx]"
```
## Development Environment Setup
+9 -49
View File
@@ -23,61 +23,21 @@ Also optionally install flash-attn:
uv pip install flash-attn --no-build-isolation -v
```
## Basic Usage
## Choose a maintained recipe
### Text-to-Video Generation
The cookbook selects complete, checked-in recipes instead of mixing model,
parallelism, offload, and attention settings independently.
```python
from fastvideo import VideoGenerator
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
### Image-to-Video Generation
```python
from fastvideo import VideoGenerator, SamplingParam
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
if __name__ == '__main__':
main()
```
!!! tip "Need more control?"
Start from a maintained recipe, then use the
[configuration](../inference/configuration.md) and
[optimization](../inference/optimizations.md) guides for supported changes.
## Next Steps
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
+16
View File
@@ -64,6 +64,7 @@ column links a runnable script in `examples/inference/basic/` where one exists.
| ltx2 | `FastVideo/LTX2-Distilled-Diffusers`<br>`FastVideo/LTX2.3-Distilled-Diffusers`<br>`FastVideo/LTX-2.3-Distilled-Diffusers` | T2V | [basic_ltx2_distilled.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2_distilled.py) |
| ltx2 | `Lightricks/LTX-2.3`<br>`FastVideo/LTX2.3-base`<br>`FastVideo/LTX2.3-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
| ltx2 | `Lightricks/LTX-2`<br>`FastVideo/LTX2-base`<br>`FastVideo/LTX2-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
| mmaudio | `FastVideo/MMAudio-large-44k-v2-Diffusers` | V2A, T2A | [basic_mmaudio.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mmaudio.py) |
| matrixgame | `FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-Base-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Diffusers`<br>`mignonjia/mg_longtuning_distilled_zelda`<br>`mignonjia/mg_sf_distilled_zelda_1k_steps`<br>`mignonjia/mg_sf_distilled_zelda`<br>`mignonjia/mg_causal_zelda`<br>`mignonjia/mg_bidirectional_zelda` | I2V | [basic_matrixgame2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame2.py) |
| matrixgame | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | I2V | [basic_matrixgame3.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame3.py) |
| minimax_h3 | `MiniMaxAI/MiniMax-H3` | T2V, I2V | [T2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_t2v.py)<br>[FL2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_fl2va.py)<br>[Ref2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_ref2va.py) |
@@ -94,6 +95,10 @@ column links a runnable script in `examples/inference/basic/` where one exists.
(`StableAudioT2AConfig` / `StableAudioOpenSmallConfig`); they are registered
under the generic T2V workload option in the registry.
**Note (MMAudio)**: the registered Hugging Face model ID is reserved but not
yet public. Follow the [MMAudio inference guide](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/pipelines/basic/mmaudio/README.md)
to convert the official weights locally and set `MMAUDIO_MODEL_PATH`.
**Note (MiniMax H3)**: T2VA, FL2VA, and Ref2VA all generate video with stereo
audio. Use the Ref2VA example when passing ordered image, video, or audio
references.
@@ -173,6 +178,17 @@ optimizations: absence means **untested**, not incompatible.
| Matrix Game 3.0 Base Distilled | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | 720x1280 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
## Apple Silicon native runtime
| Release path | Model | Mode | Validated hardware | Status |
| --- | --- | --- | --- | --- |
| MLX FastWan T2V | FastWan-QAD-INT8-1.3B `[release model ID pending]` | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB unified-memory class, MLX 0.31.2 | Release candidate; requires release-owner visual sign-off |
This is a text-to-video-only source-install release. It is validated on the
hardware listed above; MLX allocator caps are not evidence of support for a
physical 16 GB Mac. See [Apple Silicon FastWan](../getting_started/installation/mps.md)
for the supported command and release gates.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
+8
View File
@@ -33,6 +33,14 @@ For the typed config/request path added during the inference API refactor:
python examples/inference/basic/basic_dmd_new_api.py
```
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
```
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
```
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
## Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
+12 -13
View File
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -12,11 +14,11 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
@@ -24,22 +26,19 @@ def main():
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
@@ -30,8 +30,7 @@ def main():
"and casting reflections onto adjacent vehicles. "
"The motion creates space in the lineup, signaling activity within the otherwise quiet station. "
"It then comes to a smooth stop, resuming its position in line. "
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
)
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene.")
generator.generate_video(
prompt,
@@ -47,4 +46,3 @@ def main():
if __name__ == "__main__":
main()
@@ -31,8 +31,7 @@ def main():
"The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. "
"The metal surface beneath the torch shows ongoing signs of heating and melting. "
"The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, "
"underscoring the ongoing nature of the welding operation."
)
"underscoring the ongoing nature of the welding operation.")
generator.generate_video(
prompt,
@@ -46,6 +45,3 @@ def main():
if __name__ == "__main__":
main()
@@ -50,4 +50,3 @@ def main():
if __name__ == "__main__":
main()
+9 -10
View File
@@ -5,6 +5,8 @@ from fastvideo import VideoGenerator
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_dmd2"
def main():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
@@ -14,10 +16,10 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
# Adjust these offload parameters if you have < 32GB of VRAM
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
@@ -25,7 +27,6 @@ def main():
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.num_frames = 81
@@ -39,18 +40,16 @@ def main():
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
start_time = time.perf_counter()
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
end_time = time.perf_counter()
gen_time2 = end_time - start_time
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate video: {gen_time} seconds")
print(f"Time taken to generate video2: {gen_time2} seconds")
+14 -20
View File
@@ -32,11 +32,9 @@ def main():
),
# PR 2 still routes a few advanced inference knobs through the
# compatibility bridge until they get first-class typed fields.
pipeline=PipelineSelection(
experimental={
"VSA_sparsity": 0.8,
},
),
pipeline=PipelineSelection(experimental={
"VSA_sparsity": 0.8,
}, ),
)
load_start_time = time.perf_counter()
@@ -44,14 +42,12 @@ def main():
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
prompt = (
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
"LED umbrella. Steam rises from a street food cart, and a cat darts "
"across the screen. Raindrops are visible on the camera lens, creating "
"a cinematic bokeh effect."
)
prompt = ("A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
"LED umbrella. Steam rises from a street food cart, and a cat darts "
"across the screen. Raindrops are visible on the camera lens, creating "
"a cinematic bokeh effect.")
request = GenerationRequest(
prompt=prompt,
output=OutputConfig(
@@ -66,13 +62,11 @@ def main():
end_time = time.perf_counter()
gen_time = end_time - start_time
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently "
"in the breeze, enhancing the lion's commanding presence. The tone is "
"vibrant, embodying the raw energy of the wild. Low angle, steady "
"tracking shot, cinematic."
)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently "
"in the breeze, enhancing the lion's commanding presence. The tone is "
"vibrant, embodying the raw energy of the wild. Low angle, steady "
"tracking shot, cinematic.")
request2 = GenerationRequest(
prompt=prompt2,
output=OutputConfig(
@@ -2,7 +2,6 @@ import os
from fastvideo import VideoGenerator
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
@@ -46,10 +45,8 @@ def main():
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
"action_speed_list": [
float(value)
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
],
"action_speed_list":
[float(value) for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")],
}
if image_path:
kwargs["image_path"] = image_path
+180
View File
@@ -0,0 +1,180 @@
# SPDX-License-Identifier: Apache-2.0
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
sampler's shift-12 schedule instead of the base model's 50 steps, generating
synchronized video and audio in one pipeline call.
The student was trained with block-sparse video attention (VSA, 64-token
tiles) and its checkpoint carries the trained sparse-gate parameters
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
dense (every tile is selected); raise the sparsity for additional speedup.
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
# The HF repo is private while the MiniMax H3 Community License review
# completes; until it flips public, pass --model-path with a local
# snapshot of the release instead (e.g. the team export at
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", default="outputs/fasth3")
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1344)
parser.add_argument("--num-frames", type=int, default=124)
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
# default here is 5. Other grids are off-distribution.
parser.add_argument("--steps",
type=int,
default=5,
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
"forwards. 5 (default) is the distilled 4-forward grid")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--vsa-sparsity",
type=float,
default=0.0,
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
"exactly dense attention; the student was trained at 0.9")
# 64 is the trained contract: the student was TRAINED with 64-token
# (4,4,4) tiles, and its to_gate_compress gates were learned against
# pooling at that granularity — keep 64 unless you are ablating.
parser.add_argument("--vsa-tile-size",
type=int,
choices=(64, 256),
default=64,
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
"geometry for ablations")
parser.add_argument("--vsa-kernel",
choices=("triton", "sm100a"),
default="triton",
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
"fastvideo-kernel build that carries the extension; if a precondition fails at "
"run time the attention layer logs one warning and falls back to Triton. Only "
"meaningful with --vsa-tile-size 64")
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--repeats",
type=int,
default=1,
help="generate N times; with --torch-compile the first run pays "
"compilation, so steady-state is the last repeat")
return parser.parse_args()
def main() -> None:
args = parse_args()
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
if args.vsa_kernel == "sm100a":
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
# before the pipeline boots so spawned GPU workers inherit it. The
# kernel is forward-only and inference runs under no-grad, so every
# denoising forward qualifies for the CUDA route.
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
# Boot-time run configuration, folded into FastVideoArgs (the same route
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
# - attention_backend: the checkpoint carries trained to_gate_compress
# gates, which only exist under the VSA-H3 backend — a dense-backend
# load would reject them as unexpected weights. Layers that do not
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
# branch pools per tile, and the gates were trained at 64 tokens/tile.
experimental: dict[str, object] = {
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
"VSA_tile_size": args.vsa_tile_size,
}
if args.vsa_sparsity > 0.0:
experimental["VSA_sparsity"] = args.vsa_sparsity
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
pipeline=PipelineSelection(experimental=experimental),
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=args.num_gpus > 1,
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
offload=OffloadConfig(
dit=False,
dit_layerwise=False,
text_encoder=True,
vae=True,
pin_cpu_memory=False,
),
compile=CompileConfig(
enabled=args.torch_compile,
mode=args.compile_mode,
),
),
))
try:
request = GenerationRequest(
prompt=args.prompt,
negative_prompt="",
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=24,
num_inference_steps=args.steps,
# the base model is guidance-distilled; the student inherits it
guidance_scale=1.0,
batch_cfg=False,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output_dir / "fasth3.mp4"),
save_video=True,
return_frames=False,
),
)
result = generator.generate(request)
print(f"Output written to: {result.video_path}")
if result.generation_time is not None:
# machine-readable: benchmark harnesses parse this line to separate
# generation from model-load time (last occurrence = steady state)
print(f"Generation time: {result.generation_time:.2f}s")
for _ in range(args.repeats - 1):
result = generator.generate(request)
if result.generation_time is not None:
print(f"Generation time: {result.generation_time:.2f}s")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+2 -6
View File
@@ -64,12 +64,8 @@ def main() -> None:
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
tp_size = args.tp_size if args.tp_size is not None else (
args.num_gpus if args.num_gpus > 1 else 1
)
sp_size = args.sp_size if args.sp_size is not None else (
1 if args.num_gpus > 1 else args.num_gpus
)
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
generator_config = GeneratorConfig(
model_path=args.model_path,
@@ -21,7 +21,6 @@ from fastvideo.api import (
SamplingConfig,
)
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
+5 -11
View File
@@ -9,11 +9,9 @@ import re
DEFAULT_PROMPTS = [
"a photo of a cat",
(
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"
),
("a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"),
]
@@ -42,9 +40,7 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
)
p = argparse.ArgumentParser(description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.", )
p.add_argument(
"--model-path",
default="official_weights/FLUX.1-dev",
@@ -108,9 +104,7 @@ def main() -> None:
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = (
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
)
filename_base = (f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}")
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
+10 -11
View File
@@ -33,21 +33,20 @@ MODEL_PATH = os.environ.get("GAMECRAFT_MODEL_PATH", "FastVideo/HunyuanGameCraft-
# Default prompts for demo
DEFAULT_PROMPTS = {
"village": "A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
"temple": "A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
"forest": "A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
"village":
"A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
"temple":
"A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
"forest":
"A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
"beach": "A tropical beach with crystal clear turquoise water, white sand, and palm trees swaying in the breeze.",
}
# I2V: default reference image (URL). Can override with a local path.
DEFAULT_I2V_IMAGE_URL = (
"https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg"
)
DEFAULT_I2V_PROMPT = (
"An astronaut hatching from an egg, on the surface of the moon, "
"the darkness and depth of space realised in the background."
)
DEFAULT_I2V_IMAGE_URL = ("https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg")
DEFAULT_I2V_PROMPT = ("An astronaut hatching from an egg, on the surface of the moon, "
"the darkness and depth of space realised in the background.")
OUTPUT_PATH = "video_samples_gamecraft"
+16 -32
View File
@@ -26,51 +26,35 @@ from fastvideo import VideoGenerator
def main():
parser = argparse.ArgumentParser(description="GEN3C video generation")
parser.add_argument("--model_path",
type=str,
default="converted_weights/GEN3C-Cosmos-7B")
parser.add_argument("--image_path",
type=str,
default=None,
help="Input image for 3D cache conditioning")
parser.add_argument("--prompt",
type=str,
default="A slow camera pan over a sunlit landscape.")
parser.add_argument("--model_path", type=str, default="converted_weights/GEN3C-Cosmos-7B")
parser.add_argument("--image_path", type=str, default=None, help="Input image for 3D cache conditioning")
parser.add_argument("--prompt", type=str, default="A slow camera pan over a sunlit landscape.")
parser.add_argument(
"--negative_prompt",
type=str,
default=(
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
"flickering. Overall, the video is of poor quality."
),
default=("The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
"flickering. Overall, the video is of poor quality."),
)
parser.add_argument("--trajectory",
type=str,
default="left",
choices=[
"left", "right", "up", "down", "zoom_in",
"zoom_out", "clockwise", "counterclockwise", "none"
])
parser.add_argument(
"--trajectory",
type=str,
default="left",
choices=["left", "right", "up", "down", "zoom_in", "zoom_out", "clockwise", "counterclockwise", "none"])
parser.add_argument("--movement_distance", type=float, default=0.3)
parser.add_argument("--camera_rotation",
type=str,
default="center_facing",
choices=[
"center_facing", "no_rotation",
"trajectory_aligned"
])
choices=["center_facing", "no_rotation", "trajectory_aligned"])
parser.add_argument("--height", type=int, default=704)
parser.add_argument("--width", type=int, default=1280)
parser.add_argument("--num_frames", type=int, default=121)
parser.add_argument("--num_inference_steps", type=int, default=35)
parser.add_argument("--guidance_scale", type=float, default=1.0)
parser.add_argument("--output_path",
type=str,
default="outputs_video/gen3c.mp4")
parser.add_argument("--output_path", type=str, default="outputs_video/gen3c.mp4")
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
+25 -16
View File
@@ -3,6 +3,8 @@ import json
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_hy15"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -12,31 +14,38 @@ def main():
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
video = generator.generate_video(prompt,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt="",
num_frames=81,
fps=16)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
video2 = generator.generate_video(prompt2,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt="",
num_frames=81,
fps=16)
if __name__ == "__main__":
main()
main()
+13 -14
View File
@@ -3,38 +3,37 @@ import json
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_hy15_1080p"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
@@ -6,6 +6,8 @@ DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a c
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
OUTPUT_PATH = "video_samples_hyworld"
def main():
import argparse
@@ -7,7 +7,7 @@ IMAGE_PATH = "assets/girl.png"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
num_gpus=1,
@@ -19,9 +19,7 @@ def main():
# image_encoder_cpu_offload=False,
)
prompt = (
"A woman stands up and walks away"
)
prompt = ("A woman stands up and walks away")
_ = generator.generate_video(
prompt,
image_path=IMAGE_PATH,
@@ -2,6 +2,7 @@ from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
@@ -17,21 +18,28 @@ def main():
# image_encoder_cpu_offload=False,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
_ = generator.generate_video(prompt,
output_path=OUTPUT_PATH,
save_video=True,
height=512,
width=768,
num_frames=121)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2,
output_path=OUTPUT_PATH,
save_video=True,
height=512,
width=768,
num_frames=121)
if __name__ == "__main__":
main()
main()
@@ -6,7 +6,6 @@ from pathlib import Path
from fastvideo import VideoGenerator
REPO_ROOT = Path(__file__).resolve().parents[3]
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
@@ -3,17 +3,19 @@ from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embeddin
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_lingbotworld"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/LingBot-World-Base-Cam-Diffusers",
"FastVideo/LingBot-World-Base-Cam-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
+28 -34
View File
@@ -21,20 +21,16 @@ import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = (
"A woman sits at a wooden table by the window in a cozy café. She reaches out "
"with her right hand, picks up the white coffee cup from the saucer, and gently "
"brings it to her lips to take a sip. After drinking, she places the cup back on "
"the table and looks out the window, enjoying the peaceful atmosphere."
)
PROMPT = ("A woman sits at a wooden table by the window in a cozy café. She reaches out "
"with her right hand, picks up the white coffee cup from the saucer, and gently "
"brings it to her lips to take a sip. After drinking, she places the cup back on "
"the table and looks out the window, enjoying the peaceful atmosphere.")
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards")
# Input image path
IMAGE_PATH = "assets/girl.png"
@@ -51,20 +47,20 @@ def basic_generation():
print("=" * 60)
print("LongCat I2V: Basic Generation (50 steps, 480p)")
print("=" * 60)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-I2V-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_i2v_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -79,7 +75,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -94,11 +90,11 @@ def distill_refine_generation():
print("\n" + "=" * 60)
print("LongCat I2V: Distill + Refine Pipeline")
print("=" * 60)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-I2V-Diffusers",
num_gpus=1,
@@ -111,9 +107,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_i2v_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -128,14 +124,14 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 768p)
print("\n[Stage 2] Refinement (480p -> 768p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
@@ -143,7 +139,7 @@ def distill_refine_generation():
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
# Note: Refinement uses the T2V model (not I2V) since it's upscaling the generated video
# For BSA [4, 4, 8]: latent must be divisible by 8
@@ -163,9 +159,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_i2v_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -182,7 +178,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -192,13 +188,13 @@ def main():
print("\n" + "=" * 60)
print("LongCat Image-to-Video Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
@@ -206,5 +202,3 @@ def main():
if __name__ == "__main__":
main()
+30 -36
View File
@@ -15,22 +15,18 @@ import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = (
"In a realistic photography style, a white boy around seven or eight years old "
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
"features a green lawn and several tall trees, creating a warm and loving scene."
)
PROMPT = ("In a realistic photography style, a white boy around seven or eight years old "
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
"features a green lawn and several tall trees, creating a warm and loving scene.")
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards")
SEED = 42
@@ -44,20 +40,20 @@ def basic_generation():
print("=" * 60)
print("LongCat T2V: Basic Generation (50 steps, 480p)")
print("=" * 60)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_t2v_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -71,7 +67,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -86,11 +82,11 @@ def distill_refine_generation():
print("\n" + "=" * 60)
print("LongCat T2V: Distill + Refine Pipeline")
print("=" * 60)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
@@ -103,9 +99,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_t2v_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -119,14 +115,14 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 720p)
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
@@ -134,7 +130,7 @@ def distill_refine_generation():
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
refine_generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
@@ -151,9 +147,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_t2v_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -170,7 +166,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -180,13 +176,13 @@ def main():
print("\n" + "=" * 60)
print("LongCat Text-to-Video Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
@@ -194,5 +190,3 @@ def main():
if __name__ == "__main__":
main()
+35 -45
View File
@@ -21,21 +21,17 @@ import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = (
"A person rides a motorcycle along a long, straight road that stretches between "
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
"the motorcycle centered between the guardrails, while the scenery passes by on "
"both sides. The video captures the journey from the rider's perspective, emphasizing "
"the sense of motion and adventure."
)
PROMPT = ("A person rides a motorcycle along a long, straight road that stretches between "
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
"the motorcycle centered between the guardrails, while the scenery passes by on "
"both sides. The video captures the journey from the rider's perspective, emphasizing "
"the sense of motion and adventure.")
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards")
# Input video path
VIDEO_PATH = "assets/motorcycle.mp4"
@@ -55,27 +51,25 @@ def basic_generation():
print("=" * 60)
print("LongCat VC: Basic Generation (50 steps, 480p)")
print("=" * 60)
# Check if video exists
if not os.path.exists(VIDEO_PATH):
raise FileNotFoundError(
f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path."
)
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path.")
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-VC-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_vc_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -91,7 +85,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -106,18 +100,16 @@ def distill_refine_generation():
print("\n" + "=" * 60)
print("LongCat VC: Distill + Refine Pipeline")
print("=" * 60)
# Check if video exists
if not os.path.exists(VIDEO_PATH):
raise FileNotFoundError(
f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path."
)
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path.")
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-VC-Diffusers",
num_gpus=1,
@@ -130,9 +122,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_vc_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -148,14 +140,14 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 720p)
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
@@ -163,7 +155,7 @@ def distill_refine_generation():
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
# Note: Refinement uses the T2V model (not VC) since it's upscaling the generated video
refine_generator = VideoGenerator.from_pretrained(
@@ -181,9 +173,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_vc_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -200,7 +192,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -210,13 +202,13 @@ def main():
print("\n" + "=" * 60)
print("LongCat Video Continuation Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
@@ -224,5 +216,3 @@ def main():
if __name__ == "__main__":
main()
+12 -15
View File
@@ -1,19 +1,16 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
def main() -> None:
@@ -36,4 +33,4 @@ def main() -> None:
if __name__ == "__main__":
main()
main()
@@ -67,25 +67,19 @@ _inductor.coordinate_descent_tuning = True
_inductor.coordinate_descent_check_all_directions = True
_inductor.epilogue_fusion = False
MODEL_ID = os.path.expandvars(
os.path.expanduser(
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
)
)
OUTPUT_DIR = Path(
os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v")
)
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
"FastVideo/LTX-2.3-Distilled-Diffusers")))
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v"))
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
DEFAULT_PROMPT = (
"A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel."
)
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel.")
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
# Per-stage timing helpers --------------------------------------------------
def _print_stage_breakdown(result: dict, label: str) -> float | None:
"""Print stage execution times and return the sum, or None if missing."""
logging_info = result.get("logging_info")
@@ -114,9 +108,7 @@ def _collect_stage_times(
return
for name, metrics in stages.items():
stage_order.setdefault(name, None)
stage_times.setdefault(name, []).append(
float(metrics.get("execution_time", 0.0))
)
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
def _resolve_refine_upsampler(model_root: str) -> Path:
@@ -125,21 +117,18 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
cand = Path(model_root) / name
if (cand / "config.json").is_file():
return cand
raise FileNotFoundError(
f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`."
)
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`.")
# Main ---------------------------------------------------------------------
def main() -> None:
if not I2V_IMAGE:
raise SystemExit(
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py"
)
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py")
if not Path(I2V_IMAGE).is_file():
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
@@ -201,11 +190,13 @@ def main() -> None:
common_kwargs = dict(
prompt=PROMPT,
negative_prompt="", # distilled is CFG-free; no negative needed
guidance_scale=1.0, # CFG=1 for distilled
height=1280, width=832, # portrait runway aspect
num_frames=121, fps=24, # ~5s clip
num_inference_steps=8, # distilled denoise steps
negative_prompt="", # distilled is CFG-free; no negative needed
guidance_scale=1.0, # CFG=1 for distilled
height=1280,
width=832, # portrait runway aspect
num_frames=121,
fps=24, # ~5s clip
num_inference_steps=8, # distilled denoise steps
# i2v: anchor the input image at frame 0 with full strength.
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
# JPEG conditioning image.
@@ -251,10 +242,7 @@ def main() -> None:
**common_kwargs,
)
wall = time.perf_counter() - t0
e2e = (
result.get("e2e_latency")
if isinstance(result, dict) else None
) or wall
e2e = (result.get("e2e_latency") if isinstance(result, dict) else None) or wall
measured_secs.append(e2e)
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
if isinstance(result, dict):
@@ -266,10 +254,8 @@ def main() -> None:
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
if measured_secs:
avg = sum(measured_secs) / len(measured_secs)
print(
f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
)
print(f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
if stage_times:
print(f"average stage times over {measured_runs} measured runs:")
avg_total = 0.0
@@ -100,23 +100,14 @@ _inductor.coordinate_descent_tuning = True
_inductor.coordinate_descent_check_all_directions = True
_inductor.epilogue_fusion = False
MODEL_ID = os.path.expandvars(
os.path.expanduser(
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
)
)
OUTPUT_DIR = Path(
os.getenv(
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"
)
)
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
"FastVideo/LTX-2.3-Distilled-Diffusers")))
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"))
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
DEFAULT_PROMPT = (
"A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel."
)
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel.")
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
@@ -147,9 +138,7 @@ def _collect_stage_times(
return
for name, metrics in stages.items():
stage_order.setdefault(name, None)
stage_times.setdefault(name, []).append(
float(metrics.get("execution_time", 0.0))
)
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
def _resolve_refine_upsampler(model_root: str) -> Path:
@@ -157,20 +146,16 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
cand = Path(model_root) / name
if (cand / "config.json").is_file():
return cand
raise FileNotFoundError(
f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`."
)
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`.")
def main() -> None:
if not I2V_IMAGE:
raise SystemExit(
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/"
"basic_ltx2_3_distilled_i2v_typed.py"
)
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/"
"basic_ltx2_3_distilled_i2v_typed.py")
if not Path(I2V_IMAGE).is_file():
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
@@ -220,10 +205,9 @@ def main() -> None:
# model-specific VAE precision / decoder defaults are picked up
# the same way the legacy example's
# ``PipelineConfig.from_pretrained(model_root)`` did them.
components=ComponentConfig(
upsampler_weights=str(refine_upsampler_path),
# Distilled has no refine LoRA — omit ``lora_path``.
),
components=ComponentConfig(upsampler_weights=str(refine_upsampler_path),
# Distilled has no refine LoRA — omit ``lora_path``.
),
vae_tiling=False,
preset_overrides={
"refine": {
@@ -278,11 +262,7 @@ def main() -> None:
for w in range(warmup_runs):
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
t0 = time.perf_counter()
generator.generate(
build_request(
OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7
)
)
generator.generate(build_request(OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7))
dt = time.perf_counter() - t0
warmup_secs.append(dt)
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
@@ -291,46 +271,30 @@ def main() -> None:
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
for m in range(measured_runs):
out_path = (
OUTPUT_DIR
/ f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4"
)
print(
f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}"
)
out_path = (OUTPUT_DIR / f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4")
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
t0 = time.perf_counter()
result = generator.generate(
build_request(out_path, seed=2002 + m)
)
result = generator.generate(build_request(out_path, seed=2002 + m))
wall = time.perf_counter() - t0
# ``e2e_latency`` is currently surfaced via ``result.extra``;
# ``GenerationResult`` exposes ``generation_time`` as a
# first-class field but the LTX-2 pipeline only fills the
# legacy ``e2e_latency`` key. Prefer the explicit one, fall
# back to wall-clock.
e2e = (
result.extra.get("e2e_latency")
if hasattr(result, "extra") else None
) or wall
e2e = (result.extra.get("e2e_latency") if hasattr(result, "extra") else None) or wall
measured_secs.append(e2e)
print(
f"[measured {m + 1}/{measured_runs}] "
f"e2e={e2e:.2f}s wall={wall:.2f}s"
)
print(f"[measured {m + 1}/{measured_runs}] "
f"e2e={e2e:.2f}s wall={wall:.2f}s")
_print_stage_breakdown(result, f"measured {m + 1}")
_collect_stage_times(result, stage_times, stage_order)
print("\n=== summary ===")
print(
f"warmup wall-times: "
f"{[round(x, 1) for x in warmup_secs]}"
)
print(f"warmup wall-times: "
f"{[round(x, 1) for x in warmup_secs]}")
if measured_secs:
avg = sum(measured_secs) / len(measured_secs)
print(
f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
)
print(f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
if stage_times:
print(f"average stage times over {measured_runs} measured runs:")
avg_total = 0.0
@@ -1,21 +1,21 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
import os
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
@@ -12,15 +12,11 @@ from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
from fastvideo.utils import maybe_download_model
VALIDATION_JSON = (
Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json"
)
VALIDATION_JSON = (Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json")
# Override with a local snapshot or converted directory when needed, e.g.
# export LTX2_MODEL_PATH=/raid/$USER/hf/FastVideo/LTX2-Distilled-Diffusers
MODEL_ID = os.path.expandvars(
os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers"))
)
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers")))
OUTPUT_DIR = Path("outputs_video/ltx2_distilled_fast_profile")
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
@@ -69,9 +65,7 @@ def print_stage_breakdown(
return total
def extract_sr_forward_latency(
result: dict,
) -> tuple[float | None, list[tuple[str, float]], list[str]]:
def extract_sr_forward_latency(result: dict, ) -> tuple[float | None, list[tuple[str, float]], list[str]]:
logging_info = result.get("logging_info")
if logging_info is None:
return None, [], []
@@ -89,12 +83,8 @@ def extract_sr_forward_latency(
if sr_match_substr:
is_sr_stage = sr_match_substr in stage_name_l
else:
is_sr_stage = (
"srdenoisingstage" in stage_name_l
or "sr_denoising" in stage_name_l
or "upsample" in stage_name_l
or ("refine" in stage_name_l and "denois" in stage_name_l)
)
is_sr_stage = ("srdenoisingstage" in stage_name_l or "sr_denoising" in stage_name_l
or "upsample" in stage_name_l or ("refine" in stage_name_l and "denois" in stage_name_l))
if not is_sr_stage:
continue
exec_time = float(stage_metrics.get("execution_time", 0.0))
@@ -161,11 +151,9 @@ def resolve_refine_upsampler_path(model_root: str) -> Path:
return candidate
checked = "\n".join(f" - {candidate}" for candidate in candidates)
raise FileNotFoundError(
"Could not find an LTX2 refine upsampler directory.\n"
"Checked:\n"
f"{checked}"
)
raise FileNotFoundError("Could not find an LTX2 refine upsampler directory.\n"
"Checked:\n"
f"{checked}")
def main() -> None:
@@ -314,19 +302,15 @@ def main() -> None:
measured_times = run_times[measured_start_idx:]
avg_time = sum(measured_times) / len(measured_times)
print(
f"Average video generation time over {len(measured_times)} runs "
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_time:.2f}s"
)
print(f"Average video generation time over {len(measured_times)} runs "
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_time:.2f}s")
measured_e2e_times = e2e_times[measured_start_idx:]
avg_e2e_time = sum(measured_e2e_times) / len(measured_e2e_times)
print(
f"Average end-to-end latency over {len(measured_e2e_times)} runs "
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_e2e_time:.2f}s"
)
print(f"Average end-to-end latency over {len(measured_e2e_times)} runs "
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_e2e_time:.2f}s")
if sr_forward_times:
avg_sr_forward = sum(sr_forward_times) / len(sr_forward_times)
@@ -338,10 +322,8 @@ def main() -> None:
if non_stage_overhead_times:
avg_non_stage_overhead = sum(non_stage_overhead_times) / len(non_stage_overhead_times)
print(
"Average non-stage overhead over "
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s"
)
print("Average non-stage overhead over "
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s")
else:
print("Average non-stage overhead unavailable (no stage timings).")
finally:
+22 -12
View File
@@ -13,24 +13,34 @@ MODEL_VARIANT = "base_distilled_model"
# Variant-specific settings
VARIANT_CONFIG = {
"base_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim": 4,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
"model_path":
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim":
4,
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
},
"gta_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim": 2,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
"model_path":
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim":
2,
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
},
"templerun_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim": 7,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
"model_path":
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim":
7,
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
},
}
OUTPUT_PATH = "video_samples_matrixgame2"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -42,8 +52,8 @@ def main():
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -14,27 +14,40 @@ MODEL_VARIANT = "base_distilled_model"
# Variant-specific settings
VARIANT_CONFIG = {
"base_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim": 4,
"mode": "universal",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
"model_path":
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim":
4,
"mode":
"universal",
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
},
"gta_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim": 2,
"mode": "gta_drive",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
"model_path":
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim":
2,
"mode":
"gta_drive",
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
},
"templerun_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim": 7,
"mode": "templerun",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
"model_path":
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim":
7,
"mode":
"templerun",
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
},
}
OUTPUT_PATH = "video_samples_matrixgame2"
async def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -46,8 +59,8 @@ async def main():
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -56,11 +69,8 @@ async def main():
)
max_blocks = 50
num_frames = 597
actions = {
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
"mouse": torch.zeros((num_frames, 2))
}
num_frames = 597
actions = {"keyboard": torch.zeros((num_frames, config["keyboard_dim"])), "mouse": torch.zeros((num_frames, 2))}
grid_sizes = torch.tensor([150, 44, 80])
mode = config["mode"]
@@ -81,11 +91,11 @@ async def main():
for block_id in range(max_blocks):
print(f"\n=== Block {block_id + 1}/{max_blocks} ===")
action = await get_current_action_async(mode)
keyboard_cond, mouse_cond = expand_action_to_frames(action, 12)
await generator.step_async(keyboard_cond, mouse_cond)
if (await asyncio.to_thread(input, "\nContinue? (y/n): ")).lower() == 'n':
break
@@ -25,6 +25,13 @@ from fastvideo.pipelines.basic.minimax_h3.packing import resolve_canvas_size
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
parser.add_argument("--image", required=True, help="First-frame image path.")
parser.add_argument("--last-image", help="Optional last-frame image path.")
parser.add_argument("--output", default="outputs/minimax_h3_fl2va")
@@ -25,6 +25,13 @@ from fastvideo.pipelines.basic.minimax_h3 import MiniMaxH3Reference
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
parser.add_argument("--reference-video", required=True)
parser.add_argument("--reference-audio", help="Optional additional audio reference.")
parser.add_argument("--output", default="outputs/minimax_h3_ref2va")
@@ -22,6 +22,13 @@ from fastvideo.api import (
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", default="outputs/minimax_h3_t2v")
parser.add_argument("--height", type=int, default=768)
@@ -30,13 +37,15 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--torch-compile", action="store_true",
help="torch.compile the DiT transformer path")
parser.add_argument("--compile-mode", default=None,
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--repeats", type=int, default=1,
parser.add_argument("--repeats",
type=int,
default=1,
help="generate N times; with --torch-compile the first run pays "
"compilation, so steady-state is the last repeat")
"compilation, so steady-state is the last repeat")
return parser.parse_args()
@@ -67,24 +76,24 @@ def main() -> None:
))
try:
request = GenerationRequest(
prompt=args.prompt,
negative_prompt="",
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=24,
num_inference_steps=args.steps,
guidance_scale=1.0,
batch_cfg=False,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output_dir / "minimax_h3_t2v.mp4"),
save_video=True,
return_frames=False,
),
)
prompt=args.prompt,
negative_prompt="",
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=24,
num_inference_steps=args.steps,
guidance_scale=1.0,
batch_cfg=False,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output_dir / "minimax_h3_t2v.mp4"),
save_video=True,
return_frames=False,
),
)
result = generator.generate(request)
print(f"Output written to: {result.video_path}")
if result.generation_time is not None:
+44
View File
@@ -0,0 +1,44 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio large-44k-v2 video-to-audio example."""
import argparse
import os
from fastvideo import VideoGenerator
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--video-path", required=True)
parser.add_argument("--output-path", default="outputs_audio/mmaudio.wav")
parser.add_argument("--duration-seconds", type=float, default=8.0)
parser.add_argument("--prompt", default="")
parser.add_argument("--negative-prompt", default="music")
return parser.parse_args()
def main() -> None:
args = parse_args()
generator = VideoGenerator.from_pretrained(
os.environ.get(
"MMAUDIO_MODEL_PATH",
"converted_weights/mmaudio/large_44k_v2",
),
workload_type="v2a",
num_gpus=1,
)
result = generator.generate_video(
prompt=args.prompt,
negative_prompt=args.negative_prompt,
video_path=args.video_path,
audio_end_in_s=args.duration_seconds,
output_path=args.output_path,
save_video=True,
return_frames=False,
)
print(result["video_path"])
generator.shutdown()
if __name__ == "__main__":
main()
+17 -15
View File
@@ -1,19 +1,20 @@
from fastvideo import VideoGenerator, PipelineConfig
from fastvideo.api.sampling_param import SamplingParam
def main():
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
config.text_encoder_precisions = ["fp16"]
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
pipeline_config=config,
use_fsdp_inference=False, # Disable FSDP for MPS
dit_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
disable_autocast=False,
num_gpus=1,
use_fsdp_inference=False, # Disable FSDP for MPS
dit_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
disable_autocast=False,
num_gpus=1,
)
# Create sampling parameters with reduced number of frames
@@ -23,18 +24,19 @@ def main():
sampling_param.width = 256
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, sampling_param=sampling_param)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
if __name__ == "__main__":
main()
+11 -12
View File
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -16,27 +18,24 @@ def main():
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
distributed_executor_backend="ray",
# image_encoder_cpu_offload=False,
)
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
+3 -2
View File
@@ -7,7 +7,6 @@ import os
import re
from typing import List
DEFAULT_PROMPTS = [
"a photo of a cat",
"a cinematic photo of a red panda wearing a tiny backpack, standing on a rainy neon-lit street at night, shallow depth of field, sharp focus, 35mm, bokeh",
@@ -49,7 +48,9 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Run SD3.5 Medium text-to-image with FastVideo VideoGenerator.")
p.add_argument("--model-path", default="stabilityai/stable-diffusion-3.5-medium", help="Path to local diffusers-format SD3.5 weights directory.")
p.add_argument("--model-path",
default="stabilityai/stable-diffusion-3.5-medium",
help="Path to local diffusers-format SD3.5 weights directory.")
p.add_argument(
"--out-dir",
"--outdir",
@@ -3,6 +3,8 @@ import time
from fastvideo import VideoGenerator, SamplingParam
OUTPUT_PATH = "video_samples_causal"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -13,19 +15,18 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
text_encoder_cpu_offload=False,
dit_cpu_offload=False,
)
sampling_param = SamplingParam.from_pretrained(model_name)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
main()
@@ -5,6 +5,8 @@ import json
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -14,8 +16,8 @@ def main():
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
dit_precision="fp32",
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
@@ -37,7 +39,11 @@ def main():
for prompt_image_pair in prompt_image_pairs:
prompt = prompt_image_pair["prompt"]
image_path = prompt_image_pair["image_path"]
_ = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
_ = generator.generate_video(prompt,
image_path=image_path,
output_path=OUTPUT_PATH,
save_video=True,
sampling_param=sampling_param)
if __name__ == "__main__":
@@ -5,6 +5,8 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -14,15 +16,17 @@ def main():
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
init_weights_from_safetensors=
"/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
init_weights_from_safetensors_2=
"/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
num_frame_per_block=7,
# image_encoder_cpu_offload=False,
)
@@ -31,13 +35,11 @@ def main():
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
if __name__ == "__main__":
main()
main()
@@ -54,16 +54,15 @@ PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
# ...) 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
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.")
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,
@@ -15,19 +15,17 @@ def main() -> None:
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
# set to false if using RTX 4090
# set to false if using RTX 4090
# pin_cpu_memory=False,
)
# Generate videos with the same simple API, regardless of GPU count
# TurboDiffusion defaults: guidance_scale=1.0 and num_inference_steps=4 (from config)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(
prompt,
output_path=OUTPUT_PATH,
@@ -36,13 +34,11 @@ def main() -> None:
)
# Generate another video with a different prompt, without reloading the model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic."
)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(
prompt2,
output_path=OUTPUT_PATH,
@@ -17,11 +17,9 @@ def main() -> None:
num_gpus=2,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(
prompt,
output_path=OUTPUT_PATH,
@@ -30,13 +28,11 @@ def main() -> None:
)
# Generate another video with a different prompt, without reloading the model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic."
)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(
prompt2,
output_path=OUTPUT_PATH,
@@ -18,12 +18,13 @@ def main() -> None:
)
# Example prompt and image for I2V
prompt = ("Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside.")
prompt = (
"Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
)
# Use an example image path
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
video = generator.generate_video(
prompt,
image_path=image_path,
+25 -16
View File
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -12,8 +14,8 @@ def main():
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=2,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -25,24 +27,31 @@ def main():
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
_ = generator.generate_video(prompt,
output_path=OUTPUT_PATH,
save_video=True,
height=720,
width=1280,
num_frames=81)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2,
output_path=OUTPUT_PATH,
save_video=True,
height=720,
width=1280,
num_frames=81)
if __name__ == "__main__":
main()
main()
+13 -4
View File
@@ -4,6 +4,8 @@ from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_wan2_1_Fun"
OUTPUT_NAME = "wan2.1_test"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -14,8 +16,8 @@ def main():
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -30,7 +32,14 @@ def main():
image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
video = generator.generate_video(prompt, negative_prompt=negative_prompt, image_path=image_path, video_path=control_video_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
video = generator.generate_video(prompt,
negative_prompt=negative_prompt,
image_path=image_path,
video_path=control_video_path,
output_path=OUTPUT_PATH,
output_video_name=OUTPUT_NAME,
save_video=True)
if __name__ == "__main__":
main()
main()
+13 -4
View File
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -12,8 +14,8 @@ def main():
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -24,7 +26,14 @@ def main():
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, height=832, width=480, num_frames=81)
video = generator.generate_video(prompt,
image_path=image_path,
output_path=OUTPUT_PATH,
save_video=True,
height=832,
width=480,
num_frames=81)
if __name__ == "__main__":
main()
main()
+10 -9
View File
@@ -1,6 +1,8 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -11,11 +13,11 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
@@ -28,14 +30,13 @@ def main():
# model!
# T2V mode
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
main()
main()
+1 -3
View File
@@ -20,13 +20,11 @@ from fastvideo.api import (
SamplingConfig,
)
DEFAULT_PROMPT = (
"Young Chinese woman in red Hanfu, intricate embroidery. Impeccable makeup, red floral forehead pattern. "
"Elaborate high bun, golden phoenix headdress, red flowers, beads. Holds round folding fan with lady, trees, bird. "
"Neon lightning-bolt lamp (⚡️), bright yellow glow, above extended left palm. Soft-lit outdoor night background, "
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights."
)
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights.")
DEFAULT_REVISION = "f332072aa78be7aecdf3ee76d5c247082da564a6"
@@ -0,0 +1,46 @@
# SPDX-License-Identifier: Apache-2.0
"""Tiny MLX RIFE frame-interpolation smoke test."""
from __future__ import annotations
import argparse
import time
import numpy as np
from fastvideo.mlx_runtime.rife_interp import interpolate, load_model
def main() -> None:
parser = argparse.ArgumentParser(
description="MLX RIFE 4.25 frame interpolation smoke test."
)
parser.add_argument(
"--self-test",
action="store_true",
help="Run a tiny two-frame interpolation test.",
)
args = parser.parse_args()
if not args.self_test:
raise SystemExit("Nothing to do; pass --self-test")
frame0 = np.zeros((64, 96, 3), dtype=np.uint8)
frame1 = np.zeros((64, 96, 3), dtype=np.uint8)
frame1[:, :, 0] = 255
start = time.perf_counter()
model = load_model()
load_s = time.perf_counter() - start
start = time.perf_counter()
frames = interpolate([frame0, frame1], factor=2, model=model)
interp_s = time.perf_counter() - start
assert len(frames) == 3
assert frames[1].shape == frame0.shape
assert frames[1].dtype == np.uint8
print(
"MLX RIFE self-test passed: "
f"load_s={load_s:.3f} interp_s={interp_s:.3f} shape={frames[1].shape}"
)
if __name__ == "__main__":
main()
@@ -0,0 +1,510 @@
# SPDX-License-Identifier: Apache-2.0
"""End-to-end Wan2.2-TI2V-5B generation on Apple Silicon (MLX DiT + MLX TAEHV).
Pipeline: torch/MPS UMT5 encode (shared with 1.3B) → MLXWan22DiT 3-step DMD
(warped schedule, flow_shift=5) → MLX TAEHV decode (taew2_2.pth). Fully MLX
on the heavy DiT + decode path.
PYTHONPATH=$PWD python examples/inference/basic/mlx_wan22_generate.py \
--prompt "A red fox trotting through a snowy pine forest at golden hour" \
--output-path video_samples/demo_5b/fox_5b_mlx.mp4
Decoder backends: ``taehv`` (default, MLX, ~seconds), ``taehv-torch`` (parity),
``wan-vae`` (full AutoencoderKLWan on MPS, slow).
"""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
import numpy as np
from fastvideo.mlx_runtime.fast_spatial import DEFAULT_FAST_SPATIAL_SHARPEN
from fastvideo.mlx_runtime.frame_upsample import DEFAULT_PIXEL_UPSAMPLE_MODE, PIXEL_UPSAMPLE_MODES
from fastvideo.mlx_runtime.memory import cleanup_mlx
from fastvideo.mlx_runtime.prompt_cache import (
fingerprint_digest,
load_prompt_cache,
save_prompt_cache,
text_encoder_fingerprint,
)
from fastvideo.mlx_runtime.rife_interp import aligned_keyframe_count
FASTWAN21_MODEL_ID = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
FASTWAN22_MODEL_ID = "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_HEIGHT = 448
DEFAULT_WIDTH = 832
DEFAULT_NUM_FRAMES = 121
def _resolve_model_paths(
*,
text_encoder_root: Path | None,
dit_checkpoint: Path | None,
dit_config: Path | None,
vae_root: Path | None,
mlx_checkpoint: Path | None,
decode_backend: str,
) -> tuple[Path, Path | None, Path | None, Path | None]:
"""Download only the missing assets required by the selected Wan2.2 path."""
from huggingface_hub import snapshot_download
if text_encoder_root is None:
text_encoder_root = Path(snapshot_download(
FASTWAN21_MODEL_ID,
allow_patterns=["tokenizer/*", "text_encoder/*"],
))
if mlx_checkpoint is None and (dit_checkpoint is None or dit_config is None):
patterns = []
if dit_checkpoint is None:
patterns.append("transformer/diffusion_pytorch_model.safetensors")
if dit_config is None:
patterns.append("transformer/config.json")
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=patterns))
dit_checkpoint = dit_checkpoint or model_root / "transformer/diffusion_pytorch_model.safetensors"
dit_config = dit_config or model_root / "transformer/config.json"
if decode_backend == "wan-vae" and vae_root is None:
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=["vae/*"]))
vae_root = model_root / "vae"
return text_encoder_root, dit_checkpoint, dit_config, vae_root
def _prompt_cache_fingerprint(
*,
prompt: str,
prompt_used: str,
enhance_prompt: bool,
enhance_prompt_backend: str,
text_encoder_root: Path,
max_sequence_length: int,
dtype: str,
) -> dict[str, object]:
return {
"prompt": prompt,
"prompt_used": prompt_used,
"enhance_prompt": enhance_prompt,
"enhance_prompt_backend": enhance_prompt_backend,
"text_encoder": text_encoder_fingerprint(text_encoder_root),
"max_sequence_length": max_sequence_length,
"dtype": dtype,
}
def _default_prompt_cache_path(fingerprint: dict[str, object]) -> Path:
"""Content-addressed default cache file for a prompt fingerprint.
The Wan2.1 entrypoint caches prompt embeddings by default; this one only
did so when handed an explicit ``--prompt-embeds-cache`` path, so every 5B
run paid a full UMT5 encode (~45s on an M4 Max) even for a repeat prompt.
The fingerprint already covers everything that changes the embedding, so
hash it for the filename.
"""
digest = fingerprint_digest(fingerprint)[:32]
return Path.home() / ".cache" / "fastvideo" / "prompt_embeds" / f"wan22_{digest}.npy"
def main() -> None:
parser = argparse.ArgumentParser(
description="MLX Wan2.2-5B T2V (encode → DiT DMD → TAEHV/VAE decode)"
)
parser.add_argument(
"--prompt",
default="A red fox trotting through a snowy pine forest at golden hour, cinematic",
)
parser.add_argument(
"--output-path",
type=Path,
default=Path("video_samples/demo_5b/fox_5b_mlx.mp4"),
)
parser.add_argument(
"--text-encoder-root",
type=Path,
default=None,
help="Root with text_encoder/ + tokenizer/",
)
parser.add_argument(
"--prompt-embeds-cache",
type=Path,
default=None,
help="Explicit .npy UMT5 embedding cache file. Overrides the automatic "
"content-addressed cache (--prompt-cache).",
)
parser.add_argument(
"--prompt-cache",
action=argparse.BooleanOptionalAction,
default=True,
help="Cache prompt embeddings under ~/.cache/fastvideo/prompt_embeds so "
"repeat runs skip the text encoder entirely. Default: on.",
)
parser.add_argument(
"--text-encoder-device",
choices=("auto", "cpu", "mps"),
default="cpu",
help="Device for UMT5 encoding. CPU is safest beside the 5B MLX DiT.",
)
parser.add_argument(
"--enhance-prompt",
action="store_true",
help="Apply deterministic local cinematic prompt enrichment before UMT5.",
)
parser.add_argument(
"--enhance-prompt-backend",
choices=("template",),
default="template",
help="Prompt enrichment backend.",
)
parser.add_argument(
"--dit-checkpoint",
type=Path,
default=None,
)
parser.add_argument("--dit-config", type=Path, default=None)
parser.add_argument(
"--mlx-checkpoint",
type=Path,
default=None,
help="Pre-quantized MLX DiT checkpoint directory. Rewrapped with Wan2.2 per-token conditioning.",
)
parser.add_argument("--vae-root", type=Path, default=None)
parser.add_argument("--height", type=int, default=DEFAULT_HEIGHT)
parser.add_argument("--width", type=int, default=DEFAULT_WIDTH)
parser.add_argument(
"--num-frames",
type=int,
default=DEFAULT_NUM_FRAMES,
help="Pixel frames (121 at 24fps = 5.04 seconds)",
)
parser.add_argument("--seed", type=int, default=1234)
parser.add_argument("--renoise-seed", type=int, default=0)
parser.add_argument("--fps", type=int, default=24)
parser.add_argument("--flow-shift", type=float, default=5.0)
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
parser.add_argument(
"--no-warp",
action="store_true",
help="Disable schedule warping (debug only).",
)
parser.add_argument(
"--fast",
action="store_true",
help="Generate fewer frames then RIFE-interpolate to --num-frames.",
)
parser.add_argument("--fast-factor", type=int, default=2)
parser.add_argument("--fast-sharpen", type=float, default=0.6)
parser.add_argument(
"--fast-spatial",
action="store_true",
help="Denoise and decode at reduced spatial resolution, then resample "
"the decoded frames up to the target size.",
)
parser.add_argument("--fast-spatial-scale", type=int, default=2)
parser.add_argument(
"--fast-spatial-upsample-mode",
choices=PIXEL_UPSAMPLE_MODES,
default=DEFAULT_PIXEL_UPSAMPLE_MODE,
)
parser.add_argument("--fast-spatial-sharpen", type=float, default=DEFAULT_FAST_SPATIAL_SHARPEN)
parser.add_argument(
"--refine",
action="store_true",
help="Two-pass DMD: coarse denoise, upsample/re-noise, full-res denoise.",
)
parser.add_argument("--refine-scale", type=int, default=2)
parser.add_argument(
"--refine-upsample-mode",
choices=("bilinear", "nearest"),
default="bilinear",
)
parser.add_argument("--no-refine-add-noise", action="store_true")
parser.add_argument(
"--decode-backend",
choices=("taehv", "taehv-torch", "wan-vae"),
default="taehv",
)
parser.add_argument("--save-latents", type=Path, default=None)
parser.add_argument("--metrics-json", type=Path, default=None,
help="Write measured run metadata as JSON for reports or galleries.")
parser.add_argument(
"--compile",
action="store_true",
help="Compile the DiT forward with mx.compile; fallback to eager on failure.",
)
args = parser.parse_args()
if args.fast_factor < 2:
parser.error("--fast-factor must be at least 2")
# --fast-spatial used to be rejected here because it upsampled the completed
# 48-channel latent, which is out of distribution for the decoder and gave
# black or noisy video. The upsample now runs on decoded frames, so the
# latent never leaves the grid it was denoised on and the mode is usable.
if args.refine and args.fast_spatial:
print("[wan22] --refine takes precedence over --fast-spatial")
args.text_encoder_root, args.dit_checkpoint, args.dit_config, args.vae_root = _resolve_model_paths(
text_encoder_root=args.text_encoder_root,
dit_checkpoint=args.dit_checkpoint,
dit_config=args.dit_config,
vae_root=args.vae_root,
mlx_checkpoint=args.mlx_checkpoint,
decode_backend=args.decode_backend,
)
target_frames = args.num_frames
if args.fast:
args.num_frames = aligned_keyframe_count(target_frames, args.fast_factor)
print(
f"[wan22 fast] generating {args.num_frames} frames, "
f"RIFE {args.fast_factor}x -> {target_frames}"
)
import mlx.core as mx
import torch
from examples.inference.basic.mlx_wan_prompt_to_video import (
_postprocess_video,
encode_prompt,
make_rotary_embeddings,
)
from fastvideo.mlx_runtime.fast_spatial import plan_fast_spatial
from fastvideo.mlx_runtime.refine import (
default_refine_timesteps,
plan_refine_resolutions,
prepare_refine_latents,
)
from fastvideo.mlx_runtime.wan22 import (
mlx_wan22_dit_from_diffusers_safetensors,
mlx_wan22_dit_from_mlx_checkpoint,
)
from fastvideo.mlx_runtime.wan22_sample import build_wan22_dmd_schedule, sample_wan22_dmd
from fastvideo.mlx_runtime.wan_vae import decode_latents_to_video
if args.mlx_checkpoint is not None:
config = json.loads((args.mlx_checkpoint / "mlx_dit.json").read_text())["config"]
else:
config = json.loads(args.dit_config.read_text())
patch_size = tuple(config.get("patch_size", (1, 2, 2)))
if args.refine:
active_plan = plan_refine_resolutions(
height=args.height, width=args.width, num_frames=args.num_frames,
spatial_scale=args.refine_scale, vae_spatial_compression=16,
vae_temporal_compression=4, patch_size=patch_size, enabled=True,
)
spatial_mode = "refine"
elif args.fast_spatial:
fast_spatial_plan = plan_fast_spatial(
height=args.height, width=args.width, num_frames=args.num_frames,
spatial_scale=args.fast_spatial_scale, vae_spatial_compression=16,
vae_temporal_compression=4, patch_size=patch_size,
upsample_mode=args.fast_spatial_upsample_mode,
sharpen=args.fast_spatial_sharpen, enabled=True,
)
active_plan = fast_spatial_plan.plan
spatial_mode = "fast_spatial"
else:
active_plan = plan_refine_resolutions(
height=args.height, width=args.width, num_frames=args.num_frames,
spatial_scale=1, vae_spatial_compression=16, vae_temporal_compression=4,
patch_size=patch_size, enabled=False,
)
spatial_mode = "off"
lat_h, lat_w = active_plan.stage1_latent_height, active_plan.stage1_latent_width
lat_t = active_plan.latent_frames
in_ch = int(config["in_channels"])
print(f"[5B] latent {in_ch}x{lat_t}x{lat_h}x{lat_w}", flush=True)
total_start = time.perf_counter()
prompt_for_encode = args.prompt
enhance_backend = None
enhance_elapsed_s = 0.0
if args.enhance_prompt:
from fastvideo.mlx_runtime.prompt_enhance import enhance_prompt
enhancement = enhance_prompt(args.prompt, backend=args.enhance_prompt_backend)
prompt_for_encode = enhancement.enhanced
enhance_backend = enhancement.backend
enhance_elapsed_s = enhancement.elapsed_s
print(f"[enhance] backend={enhance_backend} in {enhance_elapsed_s:.2f}s", flush=True)
print(f"[enhance] prompt: {prompt_for_encode}", flush=True)
t0 = time.perf_counter()
prompt_cache_fingerprint = _prompt_cache_fingerprint(
prompt=args.prompt,
prompt_used=prompt_for_encode,
enhance_prompt=args.enhance_prompt,
enhance_prompt_backend=args.enhance_prompt_backend,
text_encoder_root=args.text_encoder_root,
max_sequence_length=512,
dtype="fp16",
)
prompt_cache_path = args.prompt_embeds_cache
if prompt_cache_path is None and args.prompt_cache:
prompt_cache_path = _default_prompt_cache_path(prompt_cache_fingerprint)
cached_embeds = load_prompt_cache(
prompt_cache_path,
prompt_cache_fingerprint,
)
if cached_embeds is not None:
embeds = torch.from_numpy(cached_embeds).contiguous()
else:
embeds = encode_prompt(
model_root=args.text_encoder_root,
prompt=prompt_for_encode,
max_sequence_length=512,
device_arg=args.text_encoder_device,
dtype_arg="fp16",
)
save_prompt_cache(
prompt_cache_path,
embeds.cpu().numpy(),
prompt_cache_fingerprint,
)
ehs = mx.array(embeds.numpy()).astype(mx.float16)
prompt_encode_s = time.perf_counter() - t0
print(f"[5B] prompt encoded {tuple(ehs.shape)} in {prompt_encode_s:.1f}s", flush=True)
t1 = time.perf_counter()
if args.mlx_checkpoint is not None:
dit = mlx_wan22_dit_from_mlx_checkpoint(
args.mlx_checkpoint,
compile=args.compile,
)
else:
dit = mlx_wan22_dit_from_diffusers_safetensors(
args.dit_checkpoint,
args.dit_config,
dtype="fp16",
compile=args.compile,
)
dit_load_s = time.perf_counter() - t1
print(f"[5B] DiT loaded in {dit_load_s:.1f}s", flush=True)
freqs = make_rotary_embeddings(config, latent_frames=lat_t, latent_height=lat_h, latent_width=lat_w)
gen = torch.Generator().manual_seed(args.seed)
noise = mx.array(
torch.randn(1, in_ch, lat_t, lat_h, lat_w, generator=gen, dtype=torch.float32).numpy()).astype(mx.float16)
steps = [int(s) for s in args.dmd_denoising_steps.split(",") if s.strip()]
t2 = time.perf_counter()
mx.reset_peak_memory()
latents = sample_wan22_dmd(
dit,
ehs,
noise,
freqs,
dmd_denoising_steps=steps,
flow_shift=args.flow_shift,
warp_denoising_step=not args.no_warp,
seed=args.renoise_seed,
)
if spatial_mode == "refine":
schedule, warped_steps = build_wan22_dmd_schedule(
steps, flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
)
# The grid opens at sigma == 1, where the hand-off
# `(1 - sigma) * upsampled + sigma * noise` weights stage 1 at zero and
# refine silently becomes a plain full-res run. Drop the leading
# full-noise steps so stage 1 actually reaches stage 2.
stage2_warped = default_refine_timesteps(schedule, warped_steps)
stage2_steps = steps[len(warped_steps) - len(stage2_warped):]
sigma = schedule.sigma_for(stage2_warped[0])
print(f"[5B refine] stage-2 steps={stage2_steps} sigma={sigma:.4f} "
f"(stage-1 weight {1.0 - sigma:.4f})", flush=True)
latents = prepare_refine_latents(
latents, scale=args.refine_scale, sigma=sigma,
add_noise_flag=not args.no_refine_add_noise,
upsample_mode=args.refine_upsample_mode, seed=args.renoise_seed + 1,
)
freqs_stage2 = make_rotary_embeddings(
config, latent_frames=lat_t,
latent_height=active_plan.stage2_latent_height,
latent_width=active_plan.stage2_latent_width,
)
latents = sample_wan22_dmd(
dit, ehs, latents, freqs_stage2, dmd_denoising_steps=stage2_steps,
flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
seed=args.renoise_seed + 2,
)
# spatial_mode == "fast_spatial" leaves the latents on the stage-1 grid;
# the resample happens after decode, in _postprocess_video.
denoise_s = time.perf_counter() - t2
peak = mx.get_peak_memory() / (1024**3)
print(f"[5B] denoise {len(steps)} steps in {denoise_s:.1f}s, peak {peak:.2f} GiB", flush=True)
latents_np = np.array(latents.astype(mx.float32))
if args.save_latents is not None:
args.save_latents.parent.mkdir(parents=True, exist_ok=True)
np.savez(args.save_latents, latents=latents_np, prompt=args.prompt, seed=args.seed)
print(f"[5B] wrote latents {args.save_latents}", flush=True)
if spatial_mode == "refine":
del freqs_stage2
del dit, latents, ehs, noise, freqs
cleanup_mlx()
metrics = decode_latents_to_video(
latents_np,
args.output_path,
fps=args.fps,
backend=args.decode_backend,
vae_dir=args.vae_root if args.decode_backend == "wan-vae" else None,
z_dim=in_ch,
)
# One h264 round-trip for both post-decode passes (see _postprocess_video).
rife_s = 0.0
rife_request = ({
"factor": args.fast_factor,
"target_frames": target_frames,
"sharpen": args.fast_sharpen,
} if args.fast else None)
spatial_request = fast_spatial_plan if spatial_mode == "fast_spatial" else None
if rife_request is not None or spatial_request is not None:
rife_start = time.perf_counter()
_postprocess_video(
video_path=args.output_path, fps=args.fps,
rife=rife_request, spatial=spatial_request,
)
rife_s = time.perf_counter() - rife_start
print(f"[5B] decoded via {metrics['backend']} in {metrics['decode_s']:.1f}s → {args.output_path}", flush=True)
summary = {
"output_path": str(args.output_path.resolve()),
"prompt": args.prompt,
"prompt_used": prompt_for_encode,
"enhance_prompt": args.enhance_prompt,
"enhance_backend": enhance_backend,
"enhance_elapsed_s": round(enhance_elapsed_s, 3),
"height": args.height,
"width": args.width,
"fps": args.fps,
"target_frames": target_frames,
"generated_frames": args.num_frames,
"seed": args.seed,
"renoise_seed": args.renoise_seed,
"dmd_denoising_steps": steps,
"flow_shift": args.flow_shift,
"warp": not args.no_warp,
"spatial_mode": spatial_mode,
"fast": args.fast,
"fast_factor": args.fast_factor if args.fast else None,
"fast_spatial_scale": args.fast_spatial_scale if args.fast_spatial else None,
"refine_scale": args.refine_scale if args.refine else None,
"decode_backend": args.decode_backend,
"prompt_encode_s": round(prompt_encode_s, 3),
"dit_load_s": round(dit_load_s, 3),
"denoise_s": round(denoise_s, 3),
"decode_s": round(metrics["decode_s"], 3),
"rife_s": round(rife_s, 3),
"wall_total_s": round(time.perf_counter() - total_start, 3),
"peak_gib": round(peak, 3),
"latent_shape": [in_ch, lat_t, lat_h, lat_w],
"stage2_latent_shape": [in_ch, lat_t, active_plan.stage2_latent_height, active_plan.stage2_latent_width],
"mlx_checkpoint": str(args.mlx_checkpoint.resolve()) if args.mlx_checkpoint else None,
}
if args.metrics_json is not None:
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
args.metrics_json.write_text(json.dumps(summary, indent=2) + "\n")
print(f"[5B] wrote metrics {args.metrics_json}", flush=True)
print(json.dumps(summary, indent=2), flush=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,130 @@
"""Compare Wan VAE and TAEHV decode on saved FastWan latents."""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import DEFAULT_MODEL_ROOT, decode_latents_to_video
def _torch_mps_memory() -> dict[str, int | None]:
"""
Report current and recommended memory usage for the MPS backend.
Returns:
dict[str, int | None]: Memory metrics in bytes, or `None` values when
PyTorch or MPS is unavailable.
"""
try:
import torch
except ImportError:
return {
"current_allocated_bytes": None,
"driver_allocated_bytes": None,
"recommended_max_bytes": None,
}
if not torch.backends.mps.is_available():
return {
"current_allocated_bytes": None,
"driver_allocated_bytes": None,
"recommended_max_bytes": None,
}
return {
"current_allocated_bytes": int(torch.mps.current_allocated_memory()),
"driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
"recommended_max_bytes": int(torch.mps.recommended_max_memory()),
}
def _parse_backends(raw: str) -> list[str]:
"""
Parse and validate a comma-separated list of decoding backends.
Parameters:
raw (str): Comma-separated backend names.
Returns:
list[str]: Trimmed, supported backend names in input order.
Raises:
ValueError: If the input contains an unsupported backend.
"""
backends = [backend.strip() for backend in raw.split(",") if backend.strip()]
allowed = {"wan-vae", "taehv"}
unknown = sorted(set(backends) - allowed)
if unknown:
raise ValueError(f"Unsupported decode backends: {unknown}")
return backends
def main() -> None:
"""
Benchmark selected Wan latent decoding backends and record their performance metrics.
Loads the specified latent array, decodes it with each selected backend, exports the
results as MP4 files, and writes per-backend timing and Torch MPS memory metrics to
`metrics.json`.
"""
parser = argparse.ArgumentParser(description="Benchmark decode backends on saved Wan/FastWan latents.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--latents-path", type=Path, required=True)
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_decode_benchmark"))
parser.add_argument("--backends", default="wan-vae,taehv")
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--taehv-source-path", type=Path, default=None)
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
parser.add_argument("--taehv-parallel", action="store_true")
args = parser.parse_args()
latents = np.load(args.latents_path)
args.output_dir.mkdir(parents=True, exist_ok=True)
rows = []
for backend in _parse_backends(args.backends):
print(f"=== Decode backend: {backend} ===")
before = _torch_mps_memory()
start = time.perf_counter()
output_path = args.output_dir / f"{args.latents_path.stem}_{backend}.mp4"
decode_latents_to_video(
model_root=args.model_root,
latents_np=latents,
output_path=output_path,
fps=args.fps,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
backend=backend,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
)
elapsed = time.perf_counter() - start
after = _torch_mps_memory()
metrics = {
"backend": backend,
"latents_path": str(args.latents_path),
"latents_shape": list(latents.shape),
"decode_export_s": elapsed,
"torch_mps_current_before_bytes": before["current_allocated_bytes"],
"torch_mps_current_after_bytes": after["current_allocated_bytes"],
"torch_mps_driver_before_bytes": before["driver_allocated_bytes"],
"torch_mps_driver_after_bytes": after["driver_allocated_bytes"],
"torch_mps_recommended_max_bytes": after["recommended_max_bytes"],
"output_path": str(output_path),
}
rows.append(metrics)
print(json.dumps(metrics, indent=2))
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(rows, indent=2))
print(f"Wrote decode metrics to: {metrics_path}")
if __name__ == "__main__":
main()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,365 @@
"""Benchmark MLX FastWan quantization modes with one shared prompt encode."""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
from typing import cast
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import (
DEFAULT_MODEL_ROOT,
decode_latents_to_video,
encode_prompt,
make_rotary_embeddings,
)
from fastvideo.mlx_runtime.memory import cleanup_mlx
def _parse_modes(raw: str) -> list[str]:
"""
Parse and validate a comma-separated list of quantization modes.
Parameters:
raw (str): Comma-separated mode names.
Returns:
list[str]: Normalized, whitespace-trimmed mode names.
Raises:
ValueError: If any mode is unsupported.
"""
modes = [mode.strip() for mode in raw.split(",") if mode.strip()]
allowed = {"none", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"}
unknown = sorted(set(modes) - allowed)
if unknown:
raise ValueError(f"Unsupported modes: {unknown}")
return modes
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
"""
Compare candidate and baseline latent arrays using error and signal-quality metrics.
Parameters:
candidate (np.ndarray): Latent array to evaluate.
baseline (np.ndarray): Reference latent array for comparison.
Returns:
dict[str, float]: Mean squared error, mean absolute error, maximum absolute
error, and signal-to-noise ratio in decibels between the arrays.
"""
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
mse = float(np.mean(np.square(diff)))
mae = float(np.mean(np.abs(diff)))
max_abs = float(np.max(np.abs(diff)))
signal = float(np.mean(np.square(baseline.astype(np.float32))))
return {
"latent_mse_vs_fp16": mse,
"latent_mae_vs_fp16": mae,
"latent_max_abs_vs_fp16": max_abs,
"latent_snr_db_vs_fp16": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
}
def _torch_mps_memory() -> dict[str, int | None]:
"""
Report PyTorch MPS memory statistics when PyTorch MPS is available.
Returns:
dict[str, int | None]: A mapping of MPS memory metric names to byte counts, or `None` values when PyTorch or MPS is unavailable.
"""
try:
import torch
except ImportError:
return {
"torch_mps_current_allocated_bytes": None,
"torch_mps_driver_allocated_bytes": None,
"torch_mps_recommended_max_bytes": None,
}
if not torch.backends.mps.is_available():
return {
"torch_mps_current_allocated_bytes": None,
"torch_mps_driver_allocated_bytes": None,
"torch_mps_recommended_max_bytes": None,
}
return {
"torch_mps_current_allocated_bytes": int(torch.mps.current_allocated_memory()),
"torch_mps_driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
"torch_mps_recommended_max_bytes": int(torch.mps.recommended_max_memory()),
}
def _decode_with_metrics(*, args, latents: np.ndarray, output_path: Path) -> dict[str, float | int | None | str]:
"""
Decode latents to a video and collect export timing and PyTorch MPS memory metrics.
Parameters:
args: Configuration values for decoding and video export.
latents (np.ndarray): Latent representation to decode.
output_path (Path): Destination path for the exported video.
Returns:
dict[str, float | int | None | str]: Video export duration and PyTorch MPS memory measurements.
"""
before = _torch_mps_memory()
decode_start = time.perf_counter()
decode_latents_to_video(
model_root=args.model_root,
latents_np=latents,
output_path=output_path,
fps=args.fps,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
backend=args.decode_backend,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
)
decode_time = time.perf_counter() - decode_start
after = _torch_mps_memory()
return {
"decode_export_s": decode_time,
"decode_torch_mps_current_before_bytes": before["torch_mps_current_allocated_bytes"],
"decode_torch_mps_current_after_bytes": after["torch_mps_current_allocated_bytes"],
"decode_torch_mps_driver_before_bytes": before["torch_mps_driver_allocated_bytes"],
"decode_torch_mps_driver_after_bytes": after["torch_mps_driver_allocated_bytes"],
"decode_torch_mps_recommended_max_bytes": after["torch_mps_recommended_max_bytes"],
}
def _run_one_mode(
*,
mode: str,
args,
config: dict,
checkpoint_path: Path,
config_path: Path,
prompt_embeds,
freqs_cis,
):
"""
Run denoising for one quantization mode and collect performance and memory metrics.
Parameters:
mode (str): Quantization mode to benchmark.
args: Benchmark configuration, including dtype, dimensions, seed, scheduler, and denoising settings.
config (dict): Model configuration containing the input channel count.
checkpoint_path (Path): Path to the transformer checkpoint.
config_path (Path): Path to the transformer configuration.
prompt_embeds: Encoded prompt embeddings shared across benchmark modes.
freqs_cis: Rotary positional embeddings used during denoising.
Returns:
dict: The mode name, generated latent array, and metrics for model loading,
denoising, step timing, and MLX memory usage.
"""
import mlx.core as mx
import torch
from fastvideo.benchmarks.mlx_fastwan_bench import denoise_dmd_on_device
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
mx_dtype = mx.float16 if args.mlx_dtype == "fp16" else mx.float32
quantization = None if mode == "none" else mode
latent_frames = (args.num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
load_start = time.perf_counter()
mx.clear_cache()
mx.reset_peak_memory()
dit = mlx_dit_from_diffusers_safetensors(
checkpoint_path,
config_path,
dtype=args.mlx_dtype,
quantization=quantization,
)
load_time = time.perf_counter() - load_start
load_peak_memory = mx.get_peak_memory()
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
# Same torch generator sequence as the original host-round-trip loop
# (initial latents first, then one re-noise draw per intermediate step),
# so every mode still shares identical stochasticity.
generator = torch.Generator(device="cpu").manual_seed(args.seed)
latents_seed = torch.randn(
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
generator=generator,
dtype=torch.float32,
).numpy()
renoise_by_step = [
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
for _ in range(max(0, len(timesteps) - 1))
]
latents = mx.array(latents_seed).astype(mx_dtype)
encoder_hidden_states = mx.array(prompt_embeds.numpy()).astype(mx_dtype)
denoise_start = time.perf_counter()
mx.reset_peak_memory()
latents_np, step_times = denoise_dmd_on_device(
mx=mx,
dit=dit,
latents=latents,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis,
timesteps=timesteps,
renoise_by_step=renoise_by_step,
schedule=schedule,
dmd_step=dmd_step,
mx_dtype=mx_dtype,
)
denoise_time = time.perf_counter() - denoise_start
denoise_peak_memory = mx.get_peak_memory()
active_memory = mx.get_active_memory()
return {
"mode": mode,
"latents": latents_np,
"metrics": {
"mlx_dit_load_s": load_time,
"mlx_denoise_s": denoise_time,
"mlx_denoise_first_step_s": step_times[0] if step_times else None,
"mlx_load_peak_bytes": int(load_peak_memory),
"mlx_denoise_peak_bytes": int(denoise_peak_memory),
"mlx_active_after_denoise_bytes": int(active_memory),
},
}
def main() -> None:
"""
Run the MLX FastWan quantization benchmark for the selected modes and write latency, memory, output, and latent-difference metrics to the output directory.
"""
parser = argparse.ArgumentParser(description="Benchmark MLX FastWan quantization modes.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--prompt", default="A snow leopard walks across a windy mountain ridge.")
parser.add_argument("--height", type=int, default=192)
parser.add_argument("--width", type=int, default=320)
parser.add_argument("--num-frames", type=int, default=17)
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
parser.add_argument("--flow-shift", type=float, default=8.0)
parser.add_argument("--max-sequence-length", type=int, default=256)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--mlx-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--modes", default="none,int8,int4,mxfp8,mxfp4,nvfp4")
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_quant_benchmark"))
parser.add_argument("--decode-backend", choices=("none", "wan-vae", "taehv"), default="taehv")
parser.add_argument("--taehv-source-path", type=Path, default=None)
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
parser.add_argument("--taehv-parallel", action="store_true")
args = parser.parse_args()
import mlx.core as mx
import torch
mx.random.seed(args.seed)
torch.manual_seed(args.seed)
args.output_dir.mkdir(parents=True, exist_ok=True)
config_path = args.model_root / "transformer/config.json"
checkpoint_path = args.model_root / "transformer/diffusion_pytorch_model.safetensors"
config = json.loads(config_path.read_text())
latent_frames = (args.num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
prompt_start = time.perf_counter()
prompt_embeds = encode_prompt(
model_root=args.model_root,
prompt=args.prompt,
max_sequence_length=args.max_sequence_length,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
)
prompt_time = time.perf_counter() - prompt_start
freqs_cis = make_rotary_embeddings(
config,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
)
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
baseline_latents = None
rows = []
for mode in _parse_modes(args.modes):
print(f"=== MLX quant mode: {mode} ===")
mode_start = time.perf_counter()
try:
result = _run_one_mode(
mode=mode,
args=args,
config=config,
checkpoint_path=checkpoint_path,
config_path=config_path,
prompt_embeds=prompt_embeds,
freqs_cis=freqs_cis,
)
except UnsupportedMLXQuantizationError as exc:
print(f"skipping mode (unsupported by this MLX build): {exc}")
rows.append({"mode": mode, "status": "unsupported_by_mlx", "error": str(exc)})
continue
cleanup_mlx(mx)
latents = result["latents"]
if baseline_latents is None:
baseline_latents = latents
latent_path = args.output_dir / f"latents_{mode}.npy"
np.save(latent_path, latents)
decode_time = 0.0
decode_metrics = {}
output_path = None
if args.decode_backend != "none":
output_path = args.output_dir / f"video_{mode}_{args.decode_backend}_{args.height}x{args.width}x{args.num_frames}.mp4"
decode_metrics = _decode_with_metrics(args=args, latents=latents, output_path=output_path)
decode_time = cast(float, decode_metrics["decode_export_s"])
mode_total = time.perf_counter() - mode_start
mlx_denoise_peak_bytes = int(result["metrics"]["mlx_denoise_peak_bytes"])
mlx_active_bytes = int(result["metrics"]["mlx_active_after_denoise_bytes"])
metrics = {
"mode": mode,
"status": "ok",
"prompt_encode_shared_s": prompt_time,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"decode_backend": args.decode_backend,
"decode_export_s": decode_time,
"mode_total_excluding_shared_prompt_s": mode_total,
"mode_total_including_shared_prompt_s": mode_total + prompt_time,
"latents_path": str(latent_path),
"output_path": str(output_path) if output_path else None,
"mlx_denoise_peak_gib": mlx_denoise_peak_bytes / (1024**3),
"mlx_active_after_denoise_gib": mlx_active_bytes / (1024**3),
"mlx_dit_peak_under_16gb": mlx_denoise_peak_bytes < 16 * 1024**3,
"mlx_dit_active_under_16gb": mlx_active_bytes < 16 * 1024**3,
"mac_16gb_status": (
"dit_memory_fits_16gb_measured_decode_separately"
if mlx_denoise_peak_bytes < 16 * 1024**3 else "dit_memory_exceeds_16gb"
),
**result["metrics"],
**decode_metrics,
**_latent_delta_metrics(latents, baseline_latents),
}
rows.append(metrics)
print(json.dumps(metrics, indent=2))
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(rows, indent=2))
print(f"Wrote benchmark metrics to: {metrics_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,104 @@
"""Compare generated MP4s against a reference MP4 with simple pixel metrics."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import numpy as np
def _read_video(path: Path) -> np.ndarray:
"""
Read all frames from a video file as an RGB NumPy array.
Parameters:
path (Path): Path to the video file.
Returns:
np.ndarray: Video frames stacked along the first axis.
Raises:
ValueError: If the video contains no readable frames.
"""
import cv2
cap = cv2.VideoCapture(str(path))
frames = []
try:
while True:
ok, frame_bgr = cap.read()
if not ok:
break
frame_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)
frames.append(frame_rgb)
finally:
cap.release()
if not frames:
raise ValueError(f"No frames read from {path}")
return np.stack(frames, axis=0)
def _metrics(candidate: np.ndarray, reference: np.ndarray) -> dict[str, float | int | list[int]]:
"""
Compute pixel-level comparison metrics between candidate and reference video frames.
Parameters:
candidate (np.ndarray): Candidate video frames in frame, height, width, and channel order.
reference (np.ndarray): Reference video frames with the same shape as the candidate.
Returns:
dict[str, float | int | list[int]]: Frame dimensions and pixel comparison metrics, including MSE, MAE, maximum absolute difference, and PSNR in decibels.
Raises:
ValueError: If the candidate and reference arrays have different shapes.
"""
if candidate.shape != reference.shape:
raise ValueError(f"Shape mismatch: candidate={candidate.shape}, reference={reference.shape}")
candidate_f = candidate.astype(np.float32)
reference_f = reference.astype(np.float32)
diff = candidate_f - reference_f
mse = float(np.mean(np.square(diff)))
mae = float(np.mean(np.abs(diff)))
max_abs = float(np.max(np.abs(diff)))
psnr = float(20.0 * np.log10(255.0 / np.sqrt(mse))) if mse > 0 else float("inf")
return {
"frames": int(candidate.shape[0]),
"height": int(candidate.shape[1]),
"width": int(candidate.shape[2]),
"channels": int(candidate.shape[3]),
"mse_vs_reference": mse,
"mae_vs_reference": mae,
"max_abs_vs_reference": max_abs,
"psnr_db_vs_reference": psnr,
}
def main() -> None:
"""Compare candidate MP4 videos with a reference and write pixel-level metrics to a JSON file."""
parser = argparse.ArgumentParser(description="Compare MP4s against a reference MP4.")
parser.add_argument("--reference", type=Path, required=True)
parser.add_argument("--candidates", type=Path, nargs="+", required=True)
parser.add_argument("--metrics-json", type=Path, required=True)
args = parser.parse_args()
reference = _read_video(args.reference)
rows = []
for candidate_path in args.candidates:
candidate = _read_video(candidate_path)
row = {
"reference_path": str(args.reference),
"candidate_path": str(candidate_path),
**_metrics(candidate, reference),
}
rows.append(row)
print(json.dumps(row, indent=2))
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
args.metrics_json.write_text(json.dumps(rows, indent=2))
print(f"Wrote video quality metrics to: {args.metrics_json}")
if __name__ == "__main__":
main()
+7 -9
View File
@@ -35,13 +35,11 @@ import torch
# Generation — make one LTX2 video to evaluate.
# ---------------------------------------------------------------------
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\""
)
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\"")
OUTPUT_PATH = "fastvideo/tests/eval/asset/ltx2.mp4"
N_DUP = 4 # how many times to duplicate the video for the gen/ref corpora
@@ -77,6 +75,7 @@ def generate_one_ltx2_video() -> str:
# Eval — the point of the script. 4 lines from "two paths" to results.
# ---------------------------------------------------------------------
def _all_registered_metrics() -> list[str]:
"""Every metric in the registry, sorted. Combined with
``skip_missing_deps=True`` this is the "run everything that works in
@@ -103,8 +102,7 @@ def score_all_metrics(video_path: str) -> None:
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
t_init0 = time.perf_counter()
ev = create_evaluator(metrics=_all_registered_metrics(),
device="cuda:0", num_gpus=1, skip_missing_deps=True)
ev = create_evaluator(metrics=_all_registered_metrics(), device="cuda:0", num_gpus=1, skip_missing_deps=True)
t_init1 = time.perf_counter()
samples = samples_from(video=gen_dir, reference=ref_dir, text_prompt=PROMPT, fps=24.0,
extract_audio=True) # auto-extract audio track from videos
@@ -23,14 +23,12 @@ import torch
from fastvideo import VideoGenerator
from fastvideo.eval import create_evaluator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\""
)
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\"")
METRICS = [
"audio.clap_score",
+18 -20
View File
@@ -25,19 +25,17 @@ from fastvideo import VideoGenerator
from fastvideo.eval import Evaluator
from fastvideo.eval.io import build_eval_kwargs
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
# VBench sub-metrics meaningful for an arbitrary text→video sample
# (just the generated frames, optionally fps + the source prompt).
@@ -45,14 +43,14 @@ PROMPT = (
# vbench.scene, ...) are excluded — they need prompts built to a
# specific schema.
METRICS = [
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
"vbench.subject_consistency", # DINO frame-to-first cosine
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
"vbench.subject_consistency", # DINO frame-to-first cosine
"vbench.background_consistency", # DINO on background patches
"vbench.imaging_quality", # pyiqa MUSIQ
"vbench.temporal_flickering", # pixel-wise frame deltas
"vbench.motion_smoothness", # AMT frame interpolator residual
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
"vbench.imaging_quality", # pyiqa MUSIQ
"vbench.temporal_flickering", # pixel-wise frame deltas
"vbench.motion_smoothness", # AMT frame interpolator residual
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
]
+42 -33
View File
@@ -41,9 +41,8 @@ def _expected_filename(row: dict) -> str:
return row["auxiliary_info"]["expected_gen_filename"]
def _generate_videos(rows: list[dict], videos_dir: Path,
model: str, num_gpus: int,
num_frames: int, height: int, width: int) -> None:
def _generate_videos(rows: list[dict], videos_dir: Path, model: str, num_gpus: int, num_frames: int, height: int,
width: int) -> None:
from fastvideo import VideoGenerator
videos_dir.mkdir(parents=True, exist_ok=True)
@@ -59,36 +58,43 @@ def _generate_videos(rows: list[dict], videos_dir: Path,
try:
for row, out_path in todo:
gen.generate_video(
prompt=row["prompt"], output_path=str(out_path), save_video=True,
num_frames=num_frames, height=height, width=width,
prompt=row["prompt"],
output_path=str(out_path),
save_video=True,
num_frames=num_frames,
height=height,
width=width,
)
finally:
gen.shutdown()
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--dataset-root", type=Path, default=None,
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--dataset-root",
type=Path,
default=None,
help="Path to a pre-downloaded Physics-IQ release. "
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
"auto-fetching missing assets from the public bucket.")
p.add_argument("--videos-dir", type=Path,
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
"auto-fetching missing assets from the public bucket.")
p.add_argument("--videos-dir",
type=Path,
default=Path("outputs_video/bench_physics_iq"),
help="Where to read/write generated videos.")
p.add_argument("--limit", type=int, default=None,
help="Truncate to first N scenarios for smoke runs.")
p.add_argument("--limit", type=int, default=None, help="Truncate to first N scenarios for smoke runs.")
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
p.add_argument("--model",
default="Davids048/LTX2-Base-Diffusers",
help="HF repo id of the text→video generator to use.")
p.add_argument("--num-frames", type=int, default=121)
p.add_argument("--height", type=int, default=1088)
p.add_argument("--width", type=int, default=1920)
p.add_argument("--skip-generation", action="store_true",
help="Re-score existing videos under --videos-dir.")
p.add_argument("--scores-out", type=Path, default=None,
p.add_argument("--skip-generation", action="store_true", help="Re-score existing videos under --videos-dir.")
p.add_argument("--scores-out",
type=Path,
default=None,
help="Where to write per-scenario scores (JSON). "
"Defaults to <videos-dir>/scores.json.")
"Defaults to <videos-dir>/scores.json.")
args = p.parse_args()
# 1. Walk the Physics-IQ corpus. Pass --limit to the dataset
@@ -100,8 +106,13 @@ def main() -> None:
# 2. Generate (or reuse) one mp4 per scenario.
if not args.skip_generation:
_generate_videos(
rows, args.videos_dir, args.model, args.num_gpus,
args.num_frames, args.height, args.width,
rows,
args.videos_dir,
args.model,
args.num_gpus,
args.num_frames,
args.height,
args.width,
)
# 3. Score each scenario. The metric reads file paths directly out
@@ -126,28 +137,26 @@ def main() -> None:
# 4. Aggregate per the upstream scoring recipe.
metric = get_metric("physics_iq")
components = metric.aggregate_components(
[r["physics_iq"] for r in all_results]
)
components = metric.aggregate_components([r["physics_iq"] for r in all_results])
print()
print("=== Physics-IQ aggregate ===")
for name, value in components.items():
print(f" {name:24s} {value:.4f}")
detailed = [
{
"scenario": row["auxiliary_info"]["scenario_id"],
"view": row["view"],
"scenario_name": row["auxiliary_info"]["scenario_name"],
"score": results["physics_iq"].score,
}
for row, results in zip(matched, all_results)
]
detailed = [{
"scenario": row["auxiliary_info"]["scenario_id"],
"view": row["view"],
"scenario_name": row["auxiliary_info"]["scenario_name"],
"score": results["physics_iq"].score,
} for row, results in zip(matched, all_results)]
out = args.scores_out or (args.videos_dir / "scores.json")
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(
{"aggregate": components, "per_scenario": detailed},
{
"aggregate": components,
"per_scenario": detailed
},
indent=2,
))
print(f"\n[done] per-scenario scores → {out}")
+33 -27
View File
@@ -39,9 +39,8 @@ def _slugify(prompt: str, max_len: int = 100) -> str:
return re.sub(r"\s+", " ", s) or "output"
def _generate_videos(prompts: list[str], videos_dir: Path,
model: str, num_gpus: int,
num_frames: int, height: int, width: int) -> None:
def _generate_videos(prompts: list[str], videos_dir: Path, model: str, num_gpus: int, num_frames: int, height: int,
width: int) -> None:
from fastvideo import VideoGenerator
videos_dir.mkdir(parents=True, exist_ok=True)
@@ -57,53 +56,60 @@ def _generate_videos(prompts: list[str], videos_dir: Path,
try:
for prompt, out_path in todo:
gen.generate_video(
prompt=prompt, output_path=str(out_path), save_video=True,
num_frames=num_frames, height=height, width=width,
prompt=prompt,
output_path=str(out_path),
save_video=True,
num_frames=num_frames,
height=height,
width=width,
)
finally:
gen.shutdown()
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--dimensions", default="aesthetic_quality,subject_consistency",
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--dimensions",
default="aesthetic_quality,subject_consistency",
help="Comma-separated VBench dimensions (or 'all').")
p.add_argument("--limit", type=int, default=None,
help="Truncate to first N prompts for smoke runs.")
p.add_argument("--videos-dir", type=Path,
default=Path("outputs_video/bench_vbench"))
p.add_argument("--limit", type=int, default=None, help="Truncate to first N prompts for smoke runs.")
p.add_argument("--videos-dir", type=Path, default=Path("outputs_video/bench_vbench"))
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
p.add_argument("--model",
default="Davids048/LTX2-Base-Diffusers",
help="HF repo id of the text→video generator to use.")
p.add_argument("--num-frames", type=int, default=121)
p.add_argument("--height", type=int, default=1088)
p.add_argument("--width", type=int, default=1920)
p.add_argument("--fps", type=float, default=24.0,
help="Frame-rate annotation passed to fps-aware metrics.")
p.add_argument("--skip-generation", action="store_true",
p.add_argument("--fps", type=float, default=24.0, help="Frame-rate annotation passed to fps-aware metrics.")
p.add_argument("--skip-generation",
action="store_true",
help="Re-score existing videos under --videos-dir without "
"regenerating.")
p.add_argument("--scores-out", type=Path, default=None,
"regenerating.")
p.add_argument("--scores-out",
type=Path,
default=None,
help="Where to dump per-prompt scores as JSON. "
"Defaults to <videos-dir>/scores.json.")
"Defaults to <videos-dir>/scores.json.")
args = p.parse_args()
# 1. Pull prompts from VBench.
dims_arg: list[str] | str = (
args.dimensions if args.dimensions == "all"
else [d.strip() for d in args.dimensions.split(",") if d.strip()]
)
dims_arg: list[str] | str = (args.dimensions if args.dimensions == "all" else
[d.strip() for d in args.dimensions.split(",") if d.strip()])
ds = get_dataset("vbench", dimensions=dims_arg)
rows = list(ds)[: args.limit]
rows = list(ds)[:args.limit]
print(f"[load] VBench: {len(rows)} prompts across {ds.dimensions}")
# 2. Generate (or reuse) one mp4 per prompt.
if not args.skip_generation:
_generate_videos(
[row["prompt"] for row in rows],
args.videos_dir, args.model, args.num_gpus,
args.num_frames, args.height, args.width,
args.videos_dir,
args.model,
args.num_gpus,
args.num_frames,
args.height,
args.width,
)
# 3. Score each video against the requested vbench sub-metrics.
@@ -122,7 +128,7 @@ def main() -> None:
samples.append({
"video": str(video_path),
"fps": args.fps,
**row, # prompt / aux / dims
**row, # prompt / aux / dims
})
matched_rows.append(row)
+18 -9
View File
@@ -35,20 +35,29 @@ from fastvideo.eval import create_evaluator, samples_from
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--gen-dir", type=Path, required=True,
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--gen-dir",
type=Path,
required=True,
help="Directory of generated videos (.mp4, .avi, .mov, .mkv, .gif).")
p.add_argument("--reference-dir", type=Path, default=None,
p.add_argument("--reference-dir",
type=Path,
default=None,
help="Directory of reference videos. Omit to score against the cached "
"reference features (built on a previous run).")
"reference features (built on a previous run).")
p.add_argument("--device", default="cuda:0" if torch.cuda.is_available() else "cpu")
p.add_argument("--num-gpus", type=int, default=1,
p.add_argument("--num-gpus",
type=int,
default=1,
help="Number of GPU replicas. >1 fans extraction out across devices.")
p.add_argument("--cache-path", type=Path, default=None,
p.add_argument("--cache-path",
type=Path,
default=None,
help="Override the reference-feature cache path. "
"Defaults to ${FASTVIDEO_EVAL_CACHE}/fvd/real_features_i3d.pt.")
p.add_argument("--output", type=Path, default=None,
"Defaults to ${FASTVIDEO_EVAL_CACHE}/fvd/real_features_i3d.pt.")
p.add_argument("--output",
type=Path,
default=None,
help="Write the result as JSON to this path (default: stdout only).")
args = p.parse_args()
+41 -43
View File
@@ -36,57 +36,55 @@ from fastvideo import VideoGenerator
from fastvideo.eval import create_evaluator
from fastvideo.eval.io import load_video
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
DEFAULT_METRICS = [
# No-input metrics: just need the generated frames.
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
"vbench.subject_consistency", # DINO frame-to-first cosine
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
"vbench.subject_consistency", # DINO frame-to-first cosine
"vbench.background_consistency", # DINO on background patches
"vbench.imaging_quality", # pyiqa MUSIQ
"vbench.temporal_flickering", # pixel-wise frame deltas
"vbench.motion_smoothness", # AMT frame interpolator residual
"vbench.imaging_quality", # pyiqa MUSIQ
"vbench.temporal_flickering", # pixel-wise frame deltas
"vbench.motion_smoothness", # AMT frame interpolator residual
# Need fps annotation:
"vbench.dynamic_degree", # RAFT optical-flow magnitude
"vbench.dynamic_degree", # RAFT optical-flow magnitude
# Need the source prompt:
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
]
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
help="HF repo id of the LTX2 checkpoint.")
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers", help="HF repo id of the LTX2 checkpoint.")
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4",
help="Where to save the generated mp4.")
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4", help="Where to save the generated mp4.")
p.add_argument("--num-frames", type=int, default=121)
p.add_argument("--height", type=int, default=1088)
p.add_argument("--width", type=int, default=1920)
p.add_argument("--prompt", default=PROMPT)
p.add_argument("--fps", type=float, default=24.0,
p.add_argument("--fps",
type=float,
default=24.0,
help="Frame-rate annotation passed to fps-aware metrics "
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
"by default.")
p.add_argument("--metrics", default=",".join(DEFAULT_METRICS),
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
"by default.")
p.add_argument("--metrics",
default=",".join(DEFAULT_METRICS),
help="Comma-separated metric names. Pass 'all' for every "
"registered metric, or e.g. 'vbench' for the whole group.")
"registered metric, or e.g. 'vbench' for the whole group.")
p.add_argument("--scores-out", default="outputs_video/ltx2_eval/scores.json")
p.add_argument("--skip-generation", action="store_true",
p.add_argument("--skip-generation",
action="store_true",
help="Reuse an existing --output video instead of regenerating.")
return p.parse_args()
@@ -115,11 +113,10 @@ def generate(args: argparse.Namespace) -> Path:
return out
def evaluate_video(video_path: Path, prompt: str, fps: float,
metric_names) -> dict:
def evaluate_video(video_path: Path, prompt: str, fps: float, metric_names) -> dict:
print(f"[eval] loading video from {video_path}...")
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
video = video.unsqueeze(0) # → (1, T, C, H, W)
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
video = video.unsqueeze(0) # → (1, T, C, H, W)
print(f"[eval] building evaluator: {metric_names}")
evaluator = create_evaluator(metrics=metric_names, device="cuda")
@@ -132,12 +129,9 @@ def evaluate_video(video_path: Path, prompt: str, fps: float,
)
if isinstance(results, list):
results = results[0] # batch of 1
results = results[0] # batch of 1
return {
name: {"score": r.score, "details": r.details}
for name, r in results.items()
}
return {name: {"score": r.score, "details": r.details} for name, r in results.items()}
def main() -> None:
@@ -157,7 +151,11 @@ def main() -> None:
out = Path(args.scores_out)
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(
{"video": str(video_path), "prompt": args.prompt, "scores": scores},
{
"video": str(video_path),
"prompt": args.prompt,
"scores": scores
},
indent=2,
))
print(f"[done] scores written to {out}")

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