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
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
167 changed files with 23315 additions and 685 deletions
+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'
+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/).
+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
+11
View File
@@ -178,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!
+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()
@@ -27,7 +27,8 @@ def parse_args() -> argparse.Namespace:
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 converted with tools/minimax_h3/fit_adaln_basis.py).
# (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.
@@ -27,7 +27,8 @@ def parse_args() -> argparse.Namespace:
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 converted with tools/minimax_h3/fit_adaln_basis.py).
# (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.
@@ -24,7 +24,8 @@ def parse_args() -> argparse.Namespace:
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 converted with tools/minimax_h3/fit_adaln_basis.py).
# (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.
@@ -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()
+36
View File
@@ -318,6 +318,38 @@ if(BUILD_CXX_KERNELS)
# Combined FastVideo Extension
# Using name 'fastvideo_kernel_ops' to distinguish from the python package namespace
# ---------------------------------------------------------------------------
# VSA block-sparse attention forward, Blackwell (sm_100a) only.
#
# NOTE the "a" suffix: -arch=sm_100a is NOT enough -- it emits a plain sm_100 target and
# ptxas rejects every tcgen05 / setmaxnreg instruction. The explicit gencode spelling
# below is required, and matches the 10.0a entry in TORCH_CUDA_ARCH_LIST.
#
# Built for 64-token sparse blocks -- FastVideo's default (4,4,4) tiling, so no
# tile-size change and no top-k granularity change is needed. VSA_BHSD
# selects [B, H, S, D]. Both are compile-time; the Python is_supported() checks incoming
# tensors against them so callers fall back to Triton rather than getting a wrong answer.
# Read the ENVIRONMENT as well as the cache variable. When TORCH_CUDA_ARCH_LIST is
# exported (build.sh, and `pip install` with it set) the branch above only prints it --
# the cmake variable stays empty, so testing that alone silently skips the kernel and
# leaves a build that succeeds with the op missing.
set(ENABLE_VSA_SM100A OFF)
set(_VSA_ARCH_LIST "${TORCH_CUDA_ARCH_LIST}")
if(NOT _VSA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
set(_VSA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
endif()
if(_VSA_ARCH_LIST MATCHES "(^|[; ,])(10\\.0a|100a|sm_100a)([; ,]|$)")
set(ENABLE_VSA_SM100A ON)
endif()
if(ENABLE_VSA_SM100A)
message(STATUS "fastvideo-kernel: building block_sparse_sm100a (Blackwell, 64- and 128-token blocks)")
list(APPEND EXTENSION_SOURCES csrc/attention/block_sparse_sm100a.cu
csrc/attention/block_sparse_blk128_sm100a.cu)
set_source_files_properties(csrc/attention/block_sparse_sm100a.cu
csrc/attention/block_sparse_blk128_sm100a.cu PROPERTIES
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a;-DVSA_BHSD=true")
endif()
Python_add_library(fastvideo_kernel_ops MODULE USE_SABI ${SKBUILD_SABI_VERSION} WITH_SOABI
${EXTENSION_SOURCES}
)
@@ -333,10 +365,14 @@ if(BUILD_CXX_KERNELS)
# Build compile definitions list
set(COMPILE_DEFS TORCH_EXTENSION_NAME=fastvideo_kernel_ops)
if(ENABLE_VSA_SM100A)
list(APPEND COMPILE_DEFS TK_COMPILE_BLOCK_SPARSE_VSA_SM100A)
endif()
if(ENABLE_TK_KERNELS)
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
endif()
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
target_compile_options(fastvideo_kernel_ops PRIVATE
+40 -11
View File
@@ -17,7 +17,7 @@ Runtime-JIT kernels (no build step, ship in every wheel/image):
| Kernels | Where | Used when |
|---|---|---|
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| FA4 CuTe-DSL block-sparse forward/backward (VSA-128/256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
## What gets built where, and when
@@ -64,28 +64,49 @@ cd fastvideo-kernel
./build.sh --rocm
```
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
### Optional: FA4 CuTe block-sparse backend (VSA-128/256 fastpath)
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
The VSA-128/256 fastpaths (tile volume 128 or 256, on NVIDIA Blackwell / sm_100) route to the
FlashAttention-4 CuTe-DSL block-sparse kernel exposed as `flash_attn.cute`. This is
an **optional** dependency: it is imported lazily, and `video_sparse_attn`
transparently falls back to the Triton backend when it is absent (so the package is
fully usable without it).
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`,
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`
and the public/private forward-backward bridges in `flash_attn.cute.interface`) are provided upstream by
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
commit `14c377950125c70b7a9dabf9c561fca53715ac7d`, the revision FastVideo pins as
the `flash-attn-4` source in the repo-root `pyproject.toml`; other revisions may
have an incompatible `_flash_attn_fwd` signature.
have incompatible block-sparse forward/backward interfaces.
Install it under its distribution name so its own runtime stack resolves with it.
Do **not** pre-install `nvidia-cutlass-dsl` by hand: this revision pins
`nvidia-cutlass-dsl==4.6.0.dev0` exactly, and a hand-installed 4.5.x floor either
gets silently upgraded or, if something else holds it back, leaves the CuTe
kernels broken.
```bash
pip install "nvidia-cutlass-dsl>=4.5.0" torchvision
pip install "git+https://github.com/Dao-AILab/flash-attention.git@940cd9680f3315f2f06b43ab5bea2c2cf2d96806#subdirectory=flash_attn/cute"
pip install torchvision
pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@14c377950125c70b7a9dabf9c561fca53715ac7d#subdirectory=flash_attn/cute"
```
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
`tests/test_vsa256_forward*.py`.
That resolves `nvidia-cutlass-dsl` to 4.6.0.dev0 and `quack-kernels` to 0.5.3, a
combination this revision works with. A mismatched CuTe DSL only surfaces when the
kernel JIT-compiles, so the error points at CuTe internals rather than at the
install:
| Error on first VSA-128/256 CuTe call | Cause |
|---|---|
| `TypeError: fmax() missing 1 required positional argument: 'b'` | `nvidia-cutlass-dsl` 4.5.x |
| `AttributeError: module 'cutlass.cute.core' has no attribute 'ThrMma'` | `quack-kernels` older than 0.5.1 |
| `ImportError: cannot import name 'alloc_reserved_mbarrier'` | `quack-kernels` 0.6.2 or newer |
An environment whose `flash_attn.cute` came from a prebuilt flash-attn wheel rather
than from this pin hits the first row; that is what the overlay step in
`docker/Dockerfile` works around.
The CuTe kernels JIT-compile on first use. Forward and backward are verified on
Blackwell (sm_100) against `tests/test_vsa128_*.py` and `tests/test_vsa256_*.py`.
## Usage
@@ -142,6 +163,14 @@ After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
# VSA-256 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 256 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 20
# VSA-128 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 128 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 40
```
### TurboDiffusion Kernels
+48 -20
View File
@@ -2,8 +2,9 @@
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
This script benchmarks the autograd-enabled wrappers:
- 64-token TK/Triton: fastvideo_kernel.block_sparse_attn.block_sparse_attn
- 128/256-token Triton/CuTe: fastvideo_kernel.block_sparse_attn_256
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
@@ -23,9 +24,6 @@ try:
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
@@ -41,7 +39,11 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--q_seq_lens",
type=int,
nargs="+",
default=[49152],
help="Q sequence lengths (must be divisible by --block_size)")
p.add_argument("--kv_seq_lens",
type=int,
nargs="+",
@@ -51,9 +53,13 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--block_size", type=int, default=64, choices=[64, 128, 256])
p.add_argument("--force_triton",
action="store_true",
help="Force wrapper to use Triton path (if supported by shapes).")
p.add_argument("--use_cute",
action="store_true",
help="Use the optional FA4 CuTe forward/backward path (requires --block_size 128 or 256).")
return p.parse_args()
@@ -84,18 +90,38 @@ def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def _configure_backend(args: argparse.Namespace) -> None:
if args.use_cute and args.block_size not in (128, 256):
raise ValueError("--use_cute requires --block_size 128 or 256")
if args.use_cute and args.force_triton:
raise ValueError("--use_cute and --force_triton are mutually exclusive")
if args.force_triton:
os.environ.pop("FASTVIDEO_VSA_CUTEDSL", None)
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
elif args.use_cute:
os.environ.pop("FASTVIDEO_VSA_TRITON", None)
os.environ.pop("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", None)
os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1"
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
_configure_backend(args)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128, block_sparse_attn_256
bs, h, d = args.batch_size, args.num_heads, args.head_dim
block_size = args.block_size
attention = {
64: block_sparse_attn,
128: block_sparse_attn_128,
256: block_sparse_attn_256,
}[block_size]
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
@@ -105,20 +131,22 @@ def main() -> None:
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print(f"block_size={block_size}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.force_triton:
if args.use_cute:
print("dispatch: FA4 CuTe")
elif args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
if q_len % block_size != 0 or kv_len % block_size != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by {block_size}")
continue
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
num_q_blocks = q_len // block_size
num_kv_blocks = kv_len // block_size
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
@@ -129,11 +157,11 @@ def main() -> None:
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks, ), BLOCK_N, dtype=torch.int32, device="cuda")
# Variable block sizes: default full logical blocks.
variable_block_sizes = torch.full((num_kv_blocks, ), block_size, dtype=torch.int32, device="cuda")
def _fwd():
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
return attention(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
@@ -142,7 +170,7 @@ def main() -> None:
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
o_, _aux_ = attention(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
@@ -156,7 +184,7 @@ def main() -> None:
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
flops = flops_sparse_attention(bs, h, d, q_len, topk, block_size)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
@@ -0,0 +1,7 @@
// block_sparse_blk128_sm100a.cu -- the 128-token-block instantiation of the torch binding.
//
// Same source as block_sparse_sm100a.cu with VSA_BLK128 set: the kernel and launch land in
// namespace vsa_blk128 (distinct symbols, no ODR clash with the blk64 objects) and the
// exported entry point becomes block_sparse_sm100a_blk128_fwd.
#define VSA_BLK128 true
#include "block_sparse_sm100a.cu"
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,201 @@
#ifndef BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
#define BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
// Launch surface for the sm_100a VSA block-sparse FMHA forward.
//
// Everything a caller needs: a POD argument struct, a predicate saying whether this build can
// run those arguments, and one launch entry point. The benchmark in
// block_sparse_bench_sm100a.cu and the torch binding both go through here, so there is
// one tensormap construction and one launch configuration rather than two that can drift.
//
// Two compile-time knobs select the four builds:
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
// VSA_BHSD false -> [token][head][dim] (BSHD), true -> [batch][head][token][dim] (BHSD)
#include "block_sparse_kernel_sm100a.cuh"
namespace VSA_NAMESPACE {
struct BlockSparseVsaArgs {
const __nv_bfloat16* q;
const __nv_bfloat16* k;
const __nv_bfloat16* v; // natural layout; only blk128 reads it (blk64 still needs v_t)
const __nv_bfloat16* v_t; // unused: kept so the bench's V_T buffer still binds
__nv_bfloat16* o;
float* lse; // [batch, num_heads, seqlen] fp32, or nullptr
const int* q2k_idx; // [batch*num_heads*num_blocks, max_kv] int32
const int* q2k_num; // [batch*num_heads*num_blocks] int32
const int* variable_block_sizes; // [num_blocks] int32, valid tokens per block
int batch;
int num_heads;
int seqlen;
int head_dim;
int num_blocks;
int max_kv;
float sm_scale;
};
// cudaSuccess iff this build can run `a`. Deliberately conservative: the caller is expected
// to fall back to its own implementation rather than get a wrong answer.
__host__ inline cudaError_t block_sparse_supported(const BlockSparseVsaArgs& a) {
if (a.head_dim != HEAD_DIM) return cudaErrorInvalidValue; // compile-time in the kernel
if (a.num_blocks % 2 != 0) return cudaErrorInvalidValue; // a CTA owns an adjacent pair
if (a.seqlen != a.num_blocks * BLOCK) return cudaErrorInvalidValue;
if (a.max_kv < 1 || a.num_blocks < 1) return cudaErrorInvalidValue;
if (a.q == nullptr || a.k == nullptr || a.o == nullptr) return cudaErrorInvalidValue;
if (a.q2k_idx == nullptr || a.q2k_num == nullptr) return cudaErrorInvalidValue;
// FastVideo always supplies this; without it padded keys would be attended as real zeros.
if (a.variable_block_sizes == nullptr) return cudaErrorInvalidValue;
// V is read MN-major at BOTH block sizes now, so no pre-transposed V_T is ever needed.
if (a.v == nullptr) return cudaErrorInvalidValue;
return cudaSuccess;
}
__host__ inline cudaError_t launch_block_sparse_sm100a(const BlockSparseVsaArgs& a,
cudaStream_t stream) {
const cudaError_t sup = block_sparse_supported(a);
if (sup != cudaSuccess) return sup;
const int B = a.batch, H = a.num_heads, S = a.seqlen, hd = a.head_dim;
const int num_blocks = a.num_blocks, max_kv = a.max_kv;
const long tq = (long)B * S;
const int packed_mtiles_per_seq = num_blocks / 2;
const int total_work = B * H * packed_mtiles_per_seq;
constexpr bool BHSD = VSA_BHSD;
CUtensorMap tq_, tk_, tvt_, tv_, to_;
{
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16, BHSD ? (uint64_t)((long)B * H) : (uint64_t)H,
BHSD ? (uint64_t)S : (uint64_t)tq, (uint64_t)Q_SUBTILES };
uint64_t gs[3] = { BHSD ? (uint64_t)((long)S * hd) * 2u : (uint64_t)hd * 2u,
BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
(uint64_t)SUB_COLS_BF16 * 2u };
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, 1u, (uint32_t)M_TILE, (uint32_t)Q_SUBTILES };
uint32_t es[4] = { 1u, 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tq_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4,
const_cast<__nv_bfloat16*>(a.q), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
if (cuTensorMapEncodeTiled(&to_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, a.o, gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
}
{
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16,
BHSD ? (uint64_t)S : (uint64_t)tq,
BHSD ? (uint64_t)(hd / SUB_COLS_BF16)
: (uint64_t)((long)H * hd / SUB_COLS_BF16),
(uint64_t)((long)B * H) };
uint64_t gs[3] = { BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
(uint64_t)SUB_COLS_BF16 * 2u,
(uint64_t)((long)S * hd) * 2u };
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, (uint32_t)BLOCK,
BLK128 ? (uint32_t)K_SUBTILES : 1u, 1u };
uint32_t es[4] = { 1u, 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tk_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
const_cast<__nv_bfloat16*>(a.k), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
// V map is byte-for-byte the K map over a.v: MN-major V needs no transpose (blk128).
const __nv_bfloat16* vbase = a.v ? a.v : a.k;
if (cuTensorMapEncodeTiled(&tv_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
const_cast<__nv_bfloat16*>(vbase), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
}
// V_T map: blk64 only. Unused at blk128 but must still be a valid tensormap to pass by value.
{
const __nv_bfloat16* vt = a.v_t ? a.v_t : a.k;
if constexpr (BLK128) {
uint64_t gd[3] = { (uint64_t)SUB_COLS_BF16, (uint64_t)((long)H * hd),
(uint64_t)((long)tq / SUB_COLS_BF16) };
uint64_t gs[2] = { (uint64_t)tq * 2u, (uint64_t)SUB_COLS_BF16 * 2u };
uint32_t bd[3] = { (uint32_t)SUB_COLS_BF16, (uint32_t)hd, (uint32_t)V_SUBTILES };
uint32_t es[3] = { 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tvt_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 3,
const_cast<__nv_bfloat16*>(vt), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
} else {
if (make_tma_2d_tiled(&tvt_, const_cast<__nv_bfloat16*>(vt), (long)H * hd, (int)tq, hd,
SUB_COLS_BF16, 2, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
CU_TENSOR_MAP_SWIZZLE_128B) != cudaSuccess)
return cudaErrorInvalidValue;
}
}
const size_t smem =
(size_t)2 * Q_TILE_BYTES + NUM_KV_STAGES * KV_RING_SLOT_BYTES
+ (size_t)2 * M_TILE * HEAD_DIM * sizeof(__nv_bfloat16)
+ (2 * NUM_KV_STAGES + 22) * 8
+ (size_t)CLC_STAGES * (2 * 8 + 16) + 16
+ 8
+ (size_t)2 * STAT_REGIONS * STATS * sizeof(float)
+ 256;
#ifndef VSA_NAMED_BAR
#define VSA_NAMED_BAR false
#endif
#ifndef VSA_THROTTLE
#define VSA_THROTTLE false
#endif
#ifndef VSA_USE_CLC
#define VSA_USE_CLC true
#endif
constexpr bool FULL_NAMED_BAR = VSA_NAMED_BAR, EX2_EMU = true, SPLIT_P = true,
SOFTMAX_THROTTLE = VSA_THROTTLE, USE_CLC = VSA_USE_CLC,
Q_RASTER = true, MHA = true;
auto kfn = &fmha_context_bf16_gen_kernel<32, FULL_NAMED_BAR, EX2_EMU, SPLIT_P,
SOFTMAX_THROTTLE, USE_CLC, Q_RASTER, MHA,
/*RESCALE_THRESHOLD=*/8, /*BHSD=*/VSA_BHSD>;
cudaError_t e = cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
if (e != cudaSuccess) return e;
const unsigned long long magic0 = make_magic((unsigned)(H * packed_mtiles_per_seq));
const unsigned long long magic1 = make_magic((unsigned)H);
const unsigned long long magic2 = make_magic((unsigned)packed_mtiles_per_seq);
const float scale_log2 = a.sm_scale * (float)M_LOG2E;
int numSM = 0;
e = cudaDeviceGetAttribute(&numSM, cudaDevAttrMultiProcessorCount, 0);
if (e != cudaSuccess) return e;
const int num_ctas = USE_CLC ? total_work : (total_work < numSM ? total_work : numSM);
dim3 grid(num_ctas, 1, 1), block(N_WARPS * 32, 1, 1);
if (USE_CLC) {
cudaLaunchConfig_t cfg = {};
cfg.gridDim = grid; cfg.blockDim = block; cfg.dynamicSmemBytes = smem; cfg.stream = stream;
cudaLaunchAttribute cfgAttr[1];
cfgAttr[0].id = cudaLaunchAttributeClusterDimension;
cfgAttr[0].val.clusterDim.x = 1; cfgAttr[0].val.clusterDim.y = 1;
cfgAttr[0].val.clusterDim.z = 1;
cfg.attrs = cfgAttr; cfg.numAttrs = 1;
return cudaLaunchKernelEx(&cfg, kfn, tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B,
num_blocks, packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
}
kfn<<<grid, block, smem, stream>>>(tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B, num_blocks,
packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
return cudaGetLastError();
}
} // namespace VSA_NAMESPACE
// Callers (the bench, the torch binding) keep using unqualified names; each translation unit
// only ever sees the one configuration its VSA_BLK128 selected.
using namespace VSA_NAMESPACE;
#endif // BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
@@ -0,0 +1,114 @@
// block_sparse_sm100a.cu -- torch binding for the sm_100a VSA block-sparse FMHA forward.
//
// Forward only: returns (out, lse) so FastVideo's existing Triton backward keeps working
// unchanged. lse is exactly the M tensor triton_block_sparse_attn_forward writes --
// max(qk * qk_scale) + log2(l), [B, H, S] fp32 -- which is what lets
// block_sparse_attn_backward_triton run against our forward untouched.
//
// The build is fixed at compile time by two flags, so one extension carries one configuration:
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
// VSA_BHSD false -> [B, S, H, D], true -> [B, H, S, D]
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include "block_sparse_launch_sm100a.cuh"
namespace {
void check_qkv(const torch::Tensor& t, const char* name, int64_t B, int64_t H, int64_t S,
int64_t D) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.scalar_type() == at::kBFloat16, name, " must be bfloat16, got ", t.scalar_type());
TORCH_CHECK(t.dim() == 4, name, " must be 4-D, got ", t.dim(), " dims");
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
if (VSA_BHSD) {
TORCH_CHECK(t.size(0) == B && t.size(1) == H && t.size(2) == S && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", H, ",", S, ",", D, "]");
} else {
TORCH_CHECK(t.size(0) == B && t.size(1) == S && t.size(2) == H && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", S, ",", H, ",", D, "]");
}
}
void check_index(const torch::Tensor& t, const char* name) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.scalar_type() == at::kInt, name, " must be int32, got ", t.scalar_type());
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
}
} // namespace
// The exported symbol carries the block size: block_sparse_sm100a_fwd is the 64-token build,
// block_sparse_sm100a_blk128_fwd the 128-token one (block_sparse_blk128_sm100a.cu re-includes
// this file with VSA_BLK128 set). The python backend picks by the metadata's block size.
#if VSA_BLK128
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_blk128_fwd
#else
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_fwd
#endif
// Returns {out} or {out, lse}. Layout of out matches the inputs.
std::vector<torch::Tensor> BLOCK_SPARSE_SM100A_FWD(torch::Tensor q, torch::Tensor k,
torch::Tensor v,
c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx,
torch::Tensor q2k_num,
torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse) {
const at::cuda::OptionalCUDAGuard guard(device_of(q));
const int64_t B = q.size(0);
const int64_t H = VSA_BHSD ? q.size(1) : q.size(2);
const int64_t S = VSA_BHSD ? q.size(2) : q.size(1);
const int64_t D = q.size(3);
check_qkv(q, "q", B, H, S, D);
check_qkv(k, "k", B, H, S, D);
check_qkv(v, "v", B, H, S, D);
check_index(q2k_idx, "q2k_idx");
check_index(q2k_num, "q2k_num");
check_index(variable_block_sizes, "variable_block_sizes");
const int64_t num_blocks = variable_block_sizes.numel();
const int64_t max_kv = q2k_idx.size(-1);
TORCH_CHECK(S == num_blocks * BLOCK, "seqlen ", S, " must equal num_blocks (", num_blocks,
") * ", BLOCK, "; FastVideo pads the sequence up to whole blocks");
auto out = torch::empty_like(q);
torch::Tensor lse;
if (need_lse) lse = torch::empty({B, H, S}, q.options().dtype(torch::kFloat32));
BlockSparseVsaArgs a{};
a.q = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr());
a.k = reinterpret_cast<const __nv_bfloat16*>(k.data_ptr());
a.v = reinterpret_cast<const __nv_bfloat16*>(v.data_ptr());
a.v_t = v_t.has_value() ? reinterpret_cast<const __nv_bfloat16*>(v_t->data_ptr()) : nullptr;
a.o = reinterpret_cast<__nv_bfloat16*>(out.data_ptr());
a.lse = need_lse ? lse.data_ptr<float>() : nullptr;
a.q2k_idx = q2k_idx.data_ptr<int>();
a.q2k_num = q2k_num.data_ptr<int>();
a.variable_block_sizes = variable_block_sizes.data_ptr<int>();
a.batch = (int)B;
a.num_heads = (int)H;
a.seqlen = (int)S;
a.head_dim = (int)D;
a.num_blocks = (int)num_blocks;
a.max_kv = (int)max_kv;
a.sm_scale = (float)sm_scale;
// Report an unsupported regime loudly rather than returning plausible-looking wrong values.
TORCH_CHECK(block_sparse_supported(a) == cudaSuccess,
"block_sparse_sm100a: unsupported configuration -- requires head_dim==",
HEAD_DIM, ", an even num_blocks, seqlen == num_blocks*", BLOCK,
", and a variable_block_sizes tensor. Got head_dim=", D, " num_blocks=",
num_blocks, " seqlen=", S);
const cudaError_t err = launch_block_sparse_sm100a(a, at::cuda::getCurrentCUDAStream());
TORCH_CHECK(err == cudaSuccess,
"block_sparse_sm100a launch failed: ", cudaGetErrorString(err));
if (need_lse) return {out, lse};
return {out};
}
@@ -0,0 +1,877 @@
// primitives.cuh -- device primitives for the sm_100a VSA block-sparse attention
// forward: tcgen05 (alloc / mma / ld / st / commit / wait / fence), TMA load / store /
// tensormap, mbarrier, cluster launch control, setmaxnreg, fast math, and the FMHA helpers.
//
// Generated and pruned to what the kernel reaches -- do not edit by hand.
#pragma once
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cmath>
#include <cassert>
#include <cstring>
#include <vector_types.h>
#ifndef CUDA_CHECK
#define CUDA_CHECK(stmt) do { \
cudaError_t _e = (stmt); \
if (_e != cudaSuccess) { \
fprintf(stderr, "CUDA error %s:%d: %s -> %s\n", \
__FILE__, __LINE__, #stmt, cudaGetErrorString(_e)); \
std::exit(1); \
} \
} while (0)
#endif
__device__ __forceinline__
uint64_t mbarrier_arrive(uint32_t mbar_smem) {
uint64_t state;
asm volatile("mbarrier.arrive.shared::cta.b64 %0, [%1];\n"
: "=l"(state) : "r"(mbar_smem) : "memory");
return state;
}
__device__ __forceinline__
void mbarrier_arrive_nostate(uint32_t mbar_smem) {
asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n"
:: "r"(mbar_smem) : "memory");
}
__device__ __forceinline__
void mbarrier_arrive_cluster_default(uint32_t cluster_smem_addr) {
asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];\n"
:: "r"(cluster_smem_addr) : "memory");
}
__device__ __forceinline__
void mbarrier_arrive_expect_tx(uint32_t mbar_smem, uint32_t expected_bytes) {
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;\n"
:: "r"(mbar_smem), "r"(expected_bytes) : "memory");
}
__device__ __forceinline__
void mbarrier_wait_parity_suspend(uint32_t mbar_smem, uint32_t phase_parity) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1, 10000000;\n"
"@!P1 bra.uni LAB_WAIT;\n"
"}\n"
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
}
__device__ __forceinline__
void mbarrier_wait_parity(uint32_t mbar_smem, uint32_t phase_parity) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT_HOT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
"@!P1 bra.uni LAB_WAIT_HOT;\n"
"}\n"
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
}
__device__ __forceinline__
void fence_proxy_async_shared_cta() {
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__
void fence_proxy_async_shared() {
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__ void clc_try_cancel_async(
uint32_t smem_dst, uint32_t mbar_smem) {
asm volatile(
"clusterlaunchcontrol.try_cancel.async.shared::cta.mbarrier::complete_tx::bytes.b128"
" [%0], [%1];\n"
:: "r"(smem_dst), "r"(mbar_smem) : "memory");
}
__device__ __forceinline__ void clc_load_response(
uint32_t smem_slot, uint32_t& r0, uint32_t& r1,
uint32_t& r2, uint32_t& r3) {
asm volatile("ld.shared::cta.v4.b32 {%0, %1, %2, %3}, [%4];\n"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
: "r"(smem_slot));
}
template <int NUM_STAGES>
struct MbarrierPhaseTracker {
uint32_t phase[NUM_STAGES];
int idx;
__device__ __forceinline__
void init() {
for (int i = 0; i < NUM_STAGES; ++i) phase[i] = 0;
idx = 0;
}
__device__ __forceinline__
uint32_t current_phase() const { return phase[idx]; }
__device__ __forceinline__
void advance() {
phase[idx] ^= 1u;
idx = (idx + 1) % NUM_STAGES;
}
__device__ __forceinline__
int stage() const { return idx; }
};
template <int NUM_STAGES>
struct PhaseTracker {
int stage;
uint32_t phase;
__device__ __forceinline__
PhaseTracker() : stage(0), phase(0) {}
__device__ __forceinline__
void advance() {
stage++;
if (stage == NUM_STAGES) {
stage = 0;
phase ^= 1;
}
}
__device__ __forceinline__
int get_stage() const { return stage; }
__device__ __forceinline__
uint32_t get_phase() const { return phase; }
};
template <int NUM_STAGES>
struct EmptyPhaseTracker {
int stage;
uint32_t phase;
__device__ __forceinline__
EmptyPhaseTracker() : stage(0), phase(1) {}
__device__ __forceinline__
void advance() {
stage++;
if (stage == NUM_STAGES) {
stage = 0;
phase ^= 1;
}
}
__device__ __forceinline__
int get_stage() const { return stage; }
__device__ __forceinline__
uint32_t get_phase() const { return phase; }
};
template <int STAGES>
__device__ __forceinline__
void advance_stage_phase(int& stage, uint32_t& phase) {
++stage;
if (stage == STAGES) {
stage = 0;
phase ^= 1u;
}
}
static constexpr uint32_t SM100_CLC_PEER_MASK = 0xFEFFFFFF;
struct ClcTileInfo {
int m_tile;
int n_tile;
bool valid;
};
enum class ClcRasterOrder { AlongN, AlongM };
__device__ __forceinline__
void clc_arrive_expect_tx_cta(uint32_t clc_full_local_addr, uint32_t tx_bytes) {
if ((threadIdx.x & 31) == 0) {
mbarrier_arrive_expect_tx(clc_full_local_addr, tx_bytes);
}
}
__device__ __forceinline__
void clc_consumer_release(uint32_t clc_empty_local_addr) {
uint32_t peer0_addr = clc_empty_local_addr & SM100_CLC_PEER_MASK;
mbarrier_arrive_cluster_default(peer0_addr);
}
__device__ __forceinline__
void clc_consumer_release_cta(uint32_t clc_empty_local_addr) {
mbarrier_arrive_nostate(clc_empty_local_addr);
}
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER>
__device__ __forceinline__
ClcTileInfo clc_parse_response(uint32_t resp_smem_addr) {
uint32_t d0, d1, d2, d3;
fence_proxy_async_shared_cta();
clc_load_response(resp_smem_addr, d0, d1, d2, d3);
const int ctaid_x = static_cast<int>(d0);
const int ctaid_y = static_cast<int>(d1 & 0xFFFFu);
const bool valid = (d2 & 1u) != 0u;
(void)d3;
ClcTileInfo info;
info.valid = valid;
if constexpr (ORDER == ClcRasterOrder::AlongN) {
info.m_tile = ctaid_y / CLUSTER_SHAPE_M;
info.n_tile = ctaid_x / CLUSTER_SHAPE_N;
} else {
info.m_tile = ctaid_x / CLUSTER_SHAPE_M;
info.n_tile = ctaid_y / CLUSTER_SHAPE_N;
}
return info;
}
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER,
int CTA_GROUP = 2, bool SUSPEND = false>
__device__ __forceinline__
ClcTileInfo clc_fetch_next_tile(
uint64_t* clc_full_bar, uint64_t* clc_empty_bar, uint32_t* clc_response,
int clc_cons_stage, uint32_t clc_cons_phase, bool do_release) {
uint32_t full_addr = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_full_bar[clc_cons_stage]));
if constexpr (SUSPEND) mbarrier_wait_parity_suspend(full_addr, clc_cons_phase);
else mbarrier_wait_parity(full_addr, clc_cons_phase);
uint32_t resp_addr = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_response[clc_cons_stage * 4]));
ClcTileInfo t = clc_parse_response<
CLUSTER_SHAPE_M, CLUSTER_SHAPE_N, ORDER>(resp_addr);
if (do_release) {
uint32_t empty_local = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_empty_bar[clc_cons_stage]));
if constexpr (CTA_GROUP == 1) {
clc_consumer_release_cta(empty_local);
} else {
clc_consumer_release(empty_local);
}
}
return t;
}
template <int STAGES = 2>
__device__ __forceinline__
void clc_fetch_next_tile_advance(int& clc_cons_stage,
uint32_t& clc_cons_phase) {
advance_stage_phase<STAGES>(clc_cons_stage, clc_cons_phase);
}
__device__ __forceinline__ unsigned fdiv(unsigned n, unsigned long long pk) {
unsigned M = (unsigned)pk;
if (M == 0u) return n;
return __umulhi(n, M) >> (unsigned)(pk >> 32);
}
__host__ inline unsigned long long make_magic(unsigned d) {
if (d <= 1u) return 0ULL;
unsigned l = 0; while ((1u << (l + 1)) <= d) ++l;
unsigned p = 31u + l;
unsigned long long m = ((1ull << p) + (unsigned long long)d - 1ull) / d;
return (m & 0xffffffffULL) | ((unsigned long long)(p - 32u) << 32);
}
template <int BARRIER_ID>
__device__ __forceinline__
void bar_sync(uint32_t thread_count) {
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
"bar.sync: BARRIER_ID must be in [0, 15]");
asm volatile("bar.sync %0, %1;\n"
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
}
template <int BARRIER_ID>
__device__ __forceinline__
void bar_arrive(uint32_t thread_count) {
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
"bar.arrive: BARRIER_ID must be in [0, 15]");
asm volatile("bar.arrive %0, %1;\n"
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
}
__device__ __forceinline__ void full_bar_arrive(int m_tile, int band) {
switch (1 + m_tile * 4 + band) {
case 1: bar_arrive<1>(64); break;
case 2: bar_arrive<2>(64); break;
case 3: bar_arrive<3>(64); break;
case 4: bar_arrive<4>(64); break;
case 5: bar_arrive<5>(64); break;
case 6: bar_arrive<6>(64); break;
case 7: bar_arrive<7>(64); break;
case 8: bar_arrive<8>(64); break;
}
}
__device__ __forceinline__ void full_bar_wait(int m_tile, int band) {
switch (1 + m_tile * 4 + band) {
case 1: bar_sync<1>(64); break;
case 2: bar_sync<2>(64); break;
case 3: bar_sync<3>(64); break;
case 4: bar_sync<4>(64); break;
case 5: bar_sync<5>(64); break;
case 6: bar_sync<6>(64); break;
case 7: bar_sync<7>(64); break;
case 8: bar_sync<8>(64); break;
}
}
template <bool IS_CAUSAL, int K_TILE>
__device__ __forceinline__ void mask_s_row_r2p(float* scores, int k_offset, int q_pos, int seqlen_k) {
int n_keep = seqlen_k - k_offset;
if constexpr (IS_CAUSAL) {
const int causal = q_pos - k_offset + 1;
n_keep = n_keep < causal ? n_keep : causal;
}
#pragma unroll
for (int s = 0; s < K_TILE / 32; ++s) {
int m = (s + 1) * 32 - n_keep;
m = m < 0 ? 0 : (m > 32 ? 32 : m);
const uint32_t keep = (m >= 32) ? 0u : (0xFFFFFFFFu >> m);
#pragma unroll
for (int i = 0; i < 32; ++i)
if (!(keep & (1u << i))) scores[s * 32 + i] = -INFINITY;
}
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_alloc(uint32_t smem_dst_ptr,
uint32_t n_cols) {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_alloc: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
:: "r"(smem_dst_ptr), "r"(n_cols));
} else {
asm volatile(
"tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;\n"
:: "r"(smem_dst_ptr), "r"(n_cols));
}
}
__device__ __forceinline__ void tcgen05_st_32x32b_x16(
uint32_t tmem_addr, const uint32_t (&r)[16]) {
asm volatile("tcgen05.st.sync.aligned.32x32b.x16.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
"%11,%12,%13,%14,%15,%16};\n"
:: "r"(tmem_addr),
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]));
}
__device__ __forceinline__ void tcgen05_st_32x32b_x32(
uint32_t tmem_addr, const uint32_t (&r)[32]) {
asm volatile("tcgen05.st.sync.aligned.32x32b.x32.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
"%11,%12,%13,%14,%15,%16,%17,%18,%19,%20,"
"%21,%22,%23,%24,%25,%26,%27,%28,%29,%30,"
"%31,%32};\n"
:: "r"(tmem_addr),
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]),
"r"(r[16]),"r"(r[17]),"r"(r[18]),"r"(r[19]),
"r"(r[20]),"r"(r[21]),"r"(r[22]),"r"(r[23]),
"r"(r[24]),"r"(r[25]),"r"(r[26]),"r"(r[27]),
"r"(r[28]),"r"(r[29]),"r"(r[30]),"r"(r[31]));
}
__device__ __forceinline__ void tcgen05_commit1_lead(uint32_t lead, uint32_t mbar_smem_addr) {
asm volatile(
"{\n\t"
".reg .pred q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"@q tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%1];\n\t"
"}\n"
:: "r"(lead), "r"(mbar_smem_addr));
}
__device__ __forceinline__ void tcgen05_wait_st() {
asm volatile("tcgen05.wait::st.sync.aligned;\n" ::: "memory");
}
__device__ __forceinline__ void tcgen05_fence_before_thread_sync() {
asm volatile("tcgen05.fence::before_thread_sync;\n" ::: "memory");
}
__device__ __forceinline__
void tma_load_2d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int coord_x, int coord_y) {
asm volatile(
"cp.async.bulk.tensor.2d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(coord_x), "r"(coord_y)
: "memory");
}
__device__ __forceinline__
void tma_load_3d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int c0, int c1, int c2) {
asm volatile(
"cp.async.bulk.tensor.3d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2)
: "memory");
}
__device__ __forceinline__
void tma_load_4d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int c0, int c1, int c2, int c3) {
asm volatile(
"cp.async.bulk.tensor.4d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2), "r"(c3)
: "memory");
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_dealloc(uint32_t tmem_addr,
uint32_t n_cols) {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_dealloc: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols));
} else {
asm volatile(
"tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols));
}
}
__device__ __forceinline__
void tma_store_2d(const void* tensormap_ptr, int coord_x, int coord_y,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.2d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2}], [%3];\n"
:: "l"(tensormap_ptr), "r"(coord_x), "r"(coord_y),
"r"(smem_src)
: "memory");
}
__device__ __forceinline__
void tma_store_3d(const void* tensormap_ptr, int c0, int c1, int c2,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.3d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2, %3}], [%4];\n"
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(smem_src)
: "memory");
}
__device__ __forceinline__
void tma_store_4d(const void* tensormap_ptr, int c0, int c1, int c2, int c3,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.4d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2, %3, %4}], [%5];\n"
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(c3), "r"(smem_src)
: "memory");
}
inline cudaError_t make_tma_2d_tiled(
CUtensorMap* out,
const void* ptr, int rows, int cols, int box_rows, int box_cols,
int elem_bytes, CUtensorMapDataType dtype,
CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B,
CUtensorMapL2promotion l2 = CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CUtensorMapFloatOOBfill oob = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) {
uint64_t globalDim[2] = { (uint64_t)cols, (uint64_t)rows };
uint64_t globalStrides[1] = { (uint64_t)cols * (uint64_t)elem_bytes };
uint32_t boxDim[2] = { (uint32_t)box_cols, (uint32_t)box_rows };
uint32_t elemStrides[2] = { 1u, 1u };
CUresult r = cuTensorMapEncodeTiled(
out, dtype, 2,
const_cast<void*>(ptr), globalDim, globalStrides,
boxDim, elemStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle, l2, oob);
return (r == CUDA_SUCCESS) ? cudaSuccess : cudaErrorInvalidValue;
}
__device__ __forceinline__
void cp_async_bulk_commit_group() {
asm volatile("cp.async.bulk.commit_group;\n" ::: "memory");
}
template <int N>
__device__ __forceinline__
void cp_async_bulk_wait_group_read() {
asm volatile("cp.async.bulk.wait_group.read %0;\n" :: "n"(N) : "memory");
}
__device__ __forceinline__
void mbarrier_init(uint32_t mbar_smem, uint32_t arrive_count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n"
:: "r"(mbar_smem), "r"(arrive_count) : "memory");
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_relinquish_alloc_permit() {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_relinquish_alloc_permit: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;\n" ::);
} else {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::2.sync.aligned;\n" ::);
}
}
__device__ __forceinline__
void fence_mbarrier_init_release_cluster() {
asm volatile("fence.mbarrier_init.release.cluster;\n" ::: "memory");
}
__device__ __forceinline__ void tcgen05_mma_f16_ss_lead(uint32_t lead,
uint32_t tmem_c, uint64_t desc_a, uint64_t desc_b, uint32_t idesc,
bool enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p, q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"setp.ne.b32 p, %5, 0;\n\t"
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], %2, %3, %4, {%6, %7, %8, %9}, p;\n\t"
"}\n"
:: "r"(lead), "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(idesc),
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
}
__device__ __forceinline__ void tcgen05_mma_f16_ts_1sm_lead(uint32_t lead,
uint32_t tmem_c, uint32_t tmem_a, uint64_t desc_b, uint32_t idesc,
bool enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p, q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"setp.ne.b32 p, %5, 0;\n\t"
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], [%2], %3, %4, {%6, %7, %8, %9}, p;\n\t"
"}\n"
:: "r"(lead), "r"(tmem_c), "r"(tmem_a), "l"(desc_b), "r"(idesc),
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
}
enum class SmemSwizzleBlackwell : uint32_t {
None = 0,
B128_32atom = 1,
B128 = 2,
B64 = 4,
B32 = 6,
};
__device__ __host__ __forceinline__ uint64_t build_smem_desc_blackwell(
uint32_t smem_addr,
uint32_t stride_byte_offset,
uint32_t leading_byte_offset,
SmemSwizzleBlackwell swizzle = SmemSwizzleBlackwell::B128,
uint32_t base_offset = 0) {
uint64_t d = 0;
d |= static_cast<uint64_t>((smem_addr >> 4) & 0x3FFF);
d |= static_cast<uint64_t>((leading_byte_offset >> 4) & 0x3FFF) << 16;
d |= static_cast<uint64_t>((stride_byte_offset >> 4) & 0x3FFF) << 32;
d |= static_cast<uint64_t>(1) << 46;
d |= (static_cast<uint64_t>(base_offset) & 0x7) << 49;
d |= static_cast<uint64_t>(static_cast<uint32_t>(swizzle) & 0x7) << 61;
return d;
}
__device__ __forceinline__
uint32_t elect_one_sync() {
uint32_t elected;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync %0|p, 0xffffffff;\n\t"
"selp.b32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(elected));
return elected;
}
__device__ __forceinline__
uint32_t elect_one_sync(uint32_t membermask) {
uint32_t elected;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync %0|p, %1;\n\t"
"selp.b32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(elected) : "r"(membermask));
return elected;
}
template <int N>
__device__ __forceinline__
void setmaxnreg_dec() {
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
"setmaxnreg_dec: N must be in [24, 256] and a multiple of 8");
asm volatile("setmaxnreg.dec.sync.aligned.u32 %0;\n"
:: "n"(N) : "memory");
}
template <int N>
__device__ __forceinline__
void setmaxnreg_inc() {
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
"setmaxnreg_inc: N must be in [24, 256] and a multiple of 8");
asm volatile("setmaxnreg.inc.sync.aligned.u32 %0;\n"
:: "n"(N) : "memory");
}
__device__ __forceinline__
uint32_t cvt_f32x2_to_bf16x2(float a, float b) {
uint32_t r;
asm volatile("cvt.rn.bf16x2.f32 %0, %2, %1;\n"
: "=r"(r) : "f"(a), "f"(b));
return r;
}
namespace {
__device__ __forceinline__ uint64_t f32x2_bits(float2 v) {
uint64_t b; __builtin_memcpy(&b, &v, 8); return b;
}
__device__ __forceinline__ float2 f32x2_make(uint64_t b) {
float2 v; __builtin_memcpy(&v, &b, 8); return v;
}
}
__device__ __forceinline__ float2 fmul2(float2 a, float2 b) {
uint64_t d;
asm volatile("mul.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 fadd2(float2 a, float2 b) {
uint64_t d;
asm volatile("add.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 ffma2(float2 a, float2 b, float2 c) {
uint64_t d;
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
: "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)), "l"(f32x2_bits(c)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 f32x2_splat(float s) { return make_float2(s, s); }
__device__ __forceinline__ float ex2_approx_f32(float z) {
float d;
asm volatile("ex2.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(z));
return d;
}
__device__ __forceinline__ float2 ex2_emu_f32x2(float x, float y) {
uint32_t ox, oy;
asm volatile(
"{\n\t"
".reg .f32 f1,f2,f3,f4,f5,f6,f7;\n\t"
".reg .b64 l1,l2,l3,l4,l5,l6,l7,l8,l9,l10;\n\t"
".reg .s32 r1,r2,r3,r4,r5,r6,r7,r8;\n\t"
"max.f32 f1, %2, 0fC2FE0000;\n\t"
"max.f32 f2, %3, 0fC2FE0000;\n\t"
"mov.b64 l1, {f1, f2};\n\t"
"mov.f32 f3, 0f4B400000;\n\t"
"mov.b64 l2, {f3, f3};\n\t"
"add.rm.f32x2 l7, l1, l2;\n\t"
"sub.rn.f32x2 l8, l7, l2;\n\t"
"sub.rn.f32x2 l9, l1, l8;\n\t"
"mov.f32 f7, 0f3D9DF09D;\n\t"
"mov.b64 l6, {f7, f7};\n\t"
"mov.f32 f6, 0f3E6906A4;\n\t"
"mov.b64 l5, {f6, f6};\n\t"
"mov.f32 f5, 0f3F31F519;\n\t"
"mov.b64 l4, {f5, f5};\n\t"
"mov.f32 f4, 0f3F800000;\n\t"
"mov.b64 l3, {f4, f4};\n\t"
"fma.rn.f32x2 l10, l9, l6, l5;\n\t"
"fma.rn.f32x2 l10, l10, l9, l4;\n\t"
"fma.rn.f32x2 l10, l10, l9, l3;\n\t"
"mov.b64 {r1, r2}, l7;\n\t"
"mov.b64 {r3, r4}, l10;\n\t"
"shl.b32 r5, r1, 23;\n\t"
"add.s32 r7, r5, r3;\n\t"
"shl.b32 r6, r2, 23;\n\t"
"add.s32 r8, r6, r4;\n\t"
"mov.b32 %0, r7;\n\t"
"mov.b32 %1, r8;\n\t"
"}\n"
: "=r"(ox), "=r"(oy) : "f"(x), "f"(y));
float2 r; __builtin_memcpy(&r.x, &ox, 4); __builtin_memcpy(&r.y, &oy, 4);
return r;
}
__device__ __forceinline__ float rcp_approx_ftz_f32(float x) {
float d;
asm("rcp.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(x));
return d;
}
__device__ __forceinline__ uint32_t make_idesc_table44(
int M, int N,
uint32_t dtype, uint32_t atype, uint32_t btype,
bool transpose_a = false, bool transpose_b = false,
bool negate_a = false, bool negate_b = false) {
uint32_t idesc = 0;
idesc |= (dtype & 0x3) << 4;
idesc |= (atype & 0x7) << 7;
idesc |= (btype & 0x7) << 10;
idesc |= (negate_a ? 1u : 0u) << 13;
idesc |= (negate_b ? 1u : 0u) << 14;
idesc |= (transpose_a ? 1u : 0u) << 15;
idesc |= (transpose_b ? 1u : 0u) << 16;
idesc |= ((static_cast<uint32_t>(N) >> 3) & 0x3F) << 17;
idesc |= ((static_cast<uint32_t>(M) >> 4) & 0x1F) << 24;
return idesc;
}
__device__ __forceinline__ uint32_t make_idesc_bf16_f32(
int M, int N, bool ta = false, bool tb = false) {
return make_idesc_table44(M, N, 1,
1, 1, ta, tb);
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x16(
uint32_t tmem_addr, uint32_t (&r)[16]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15}, [%16];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x32(
uint32_t tmem_addr, uint32_t (&r)[32]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x32.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31}, [%32];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x64(
uint32_t tmem_addr, uint32_t (&r)[64]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x64.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
"%60,%61,%62,%63}, [%64];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31]),
"=r"(r[32]),"=r"(r[33]),"=r"(r[34]),"=r"(r[35]),
"=r"(r[36]),"=r"(r[37]),"=r"(r[38]),"=r"(r[39]),
"=r"(r[40]),"=r"(r[41]),"=r"(r[42]),"=r"(r[43]),
"=r"(r[44]),"=r"(r[45]),"=r"(r[46]),"=r"(r[47]),
"=r"(r[48]),"=r"(r[49]),"=r"(r[50]),"=r"(r[51]),
"=r"(r[52]),"=r"(r[53]),"=r"(r[54]),"=r"(r[55]),
"=r"(r[56]),"=r"(r[57]),"=r"(r[58]),"=r"(r[59]),
"=r"(r[60]),"=r"(r[61]),"=r"(r[62]),"=r"(r[63])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x128(
uint32_t tmem_addr, uint32_t (&r)[128]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x128.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
"%60,%61,%62,%63,%64,%65,%66,%67,%68,%69,"
"%70,%71,%72,%73,%74,%75,%76,%77,%78,%79,"
"%80,%81,%82,%83,%84,%85,%86,%87,%88,%89,"
"%90,%91,%92,%93,%94,%95,%96,%97,%98,%99,"
"%100,%101,%102,%103,%104,%105,%106,%107,%108,%109,"
"%110,%111,%112,%113,%114,%115,%116,%117,%118,%119,"
"%120,%121,%122,%123,%124,%125,%126,%127}, [%128];\n"
: "=r"(r[ 0]),"=r"(r[ 1]),"=r"(r[ 2]),"=r"(r[ 3]),
"=r"(r[ 4]),"=r"(r[ 5]),"=r"(r[ 6]),"=r"(r[ 7]),
"=r"(r[ 8]),"=r"(r[ 9]),"=r"(r[ 10]),"=r"(r[ 11]),
"=r"(r[ 12]),"=r"(r[ 13]),"=r"(r[ 14]),"=r"(r[ 15]),
"=r"(r[ 16]),"=r"(r[ 17]),"=r"(r[ 18]),"=r"(r[ 19]),
"=r"(r[ 20]),"=r"(r[ 21]),"=r"(r[ 22]),"=r"(r[ 23]),
"=r"(r[ 24]),"=r"(r[ 25]),"=r"(r[ 26]),"=r"(r[ 27]),
"=r"(r[ 28]),"=r"(r[ 29]),"=r"(r[ 30]),"=r"(r[ 31]),
"=r"(r[ 32]),"=r"(r[ 33]),"=r"(r[ 34]),"=r"(r[ 35]),
"=r"(r[ 36]),"=r"(r[ 37]),"=r"(r[ 38]),"=r"(r[ 39]),
"=r"(r[ 40]),"=r"(r[ 41]),"=r"(r[ 42]),"=r"(r[ 43]),
"=r"(r[ 44]),"=r"(r[ 45]),"=r"(r[ 46]),"=r"(r[ 47]),
"=r"(r[ 48]),"=r"(r[ 49]),"=r"(r[ 50]),"=r"(r[ 51]),
"=r"(r[ 52]),"=r"(r[ 53]),"=r"(r[ 54]),"=r"(r[ 55]),
"=r"(r[ 56]),"=r"(r[ 57]),"=r"(r[ 58]),"=r"(r[ 59]),
"=r"(r[ 60]),"=r"(r[ 61]),"=r"(r[ 62]),"=r"(r[ 63]),
"=r"(r[ 64]),"=r"(r[ 65]),"=r"(r[ 66]),"=r"(r[ 67]),
"=r"(r[ 68]),"=r"(r[ 69]),"=r"(r[ 70]),"=r"(r[ 71]),
"=r"(r[ 72]),"=r"(r[ 73]),"=r"(r[ 74]),"=r"(r[ 75]),
"=r"(r[ 76]),"=r"(r[ 77]),"=r"(r[ 78]),"=r"(r[ 79]),
"=r"(r[ 80]),"=r"(r[ 81]),"=r"(r[ 82]),"=r"(r[ 83]),
"=r"(r[ 84]),"=r"(r[ 85]),"=r"(r[ 86]),"=r"(r[ 87]),
"=r"(r[ 88]),"=r"(r[ 89]),"=r"(r[ 90]),"=r"(r[ 91]),
"=r"(r[ 92]),"=r"(r[ 93]),"=r"(r[ 94]),"=r"(r[ 95]),
"=r"(r[ 96]),"=r"(r[ 97]),"=r"(r[ 98]),"=r"(r[ 99]),
"=r"(r[100]),"=r"(r[101]),"=r"(r[102]),"=r"(r[103]),
"=r"(r[104]),"=r"(r[105]),"=r"(r[106]),"=r"(r[107]),
"=r"(r[108]),"=r"(r[109]),"=r"(r[110]),"=r"(r[111]),
"=r"(r[112]),"=r"(r[113]),"=r"(r[114]),"=r"(r[115]),
"=r"(r[116]),"=r"(r[117]),"=r"(r[118]),"=r"(r[119]),
"=r"(r[120]),"=r"(r[121]),"=r"(r[122]),"=r"(r[123]),
"=r"(r[124]),"=r"(r[125]),"=r"(r[126]),"=r"(r[127])
: "r"(tmem_addr));
}
__device__ __forceinline__
uint32_t smem_ptr_u32(const void* ptr) {
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
}
__device__ __forceinline__
void sts_f32(uint32_t smem_addr, float val) {
asm volatile("st.shared.f32 [%0], %1;" :: "r"(smem_addr), "f"(val) : "memory");
}
@@ -28,10 +28,31 @@ void register_rms_norm(pybind11::module_ &);
void register_layer_norm(pybind11::module_ &);
void register_gemm(pybind11::module_ &);
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
extern std::vector<torch::Tensor> block_sparse_sm100a_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse);
extern std::vector<torch::Tensor> block_sparse_sm100a_blk128_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "FastVideo CUDA Kernels";
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
m.def("block_sparse_sm100a_fwd",
torch::wrap_pybind_function(block_sparse_sm100a_fwd),
"VSA block-sparse attention forward, 64-token blocks (Blackwell sm100a)");
m.def("block_sparse_sm100a_blk128_fwd",
torch::wrap_pybind_function(block_sparse_sm100a_blk128_fwd),
"VSA block-sparse attention forward, 128-token blocks (Blackwell sm100a)");
#endif
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention (Hopper)");
#endif
@@ -1,4 +1,4 @@
"""VSA-256 block-sparse attention wrapper.
"""VSA-128/256 block-sparse attention wrappers.
The default 256-block path is Triton: it expands the logical 256-block map
to the existing 64-block Triton kernel via a dense 4x4 expansion per logical
@@ -7,8 +7,8 @@ edge ("route A"), and requires no optional dependencies.
The FA4 CuTe block-sparse fastpath (intended for Blackwell sm_100+) is
*opt-in* via ``FASTVIDEO_VSA_CUTEDSL=1``. It routes to
:mod:`fastvideo_kernel.block_sparse_attn_cute_fwd`, which natively operates
on 128-token KV blocks (this wrapper expands the logical 256-block map /
sizes into that physical 128-block representation). The CuTe kernel
on 128-token Q/KV blocks (the 256 wrapper expands its logical KV map and
sizes into that physical representation). The CuTe kernel
(``flash_attn.cute`` with block-sparsity) is an optional dependency,
imported lazily only when this fastpath is selected.
@@ -35,7 +35,7 @@ _KV_BLOCK_TRITON = 64 # Existing Triton path uses 64-token KV blocks.
def _resolve_backend() -> str:
"""Pick the backend for the 256-block VSA path.
"""Pick the backend for the 128/256-block VSA paths.
Default is Triton (no optional deps). The FA4 CuTe fastpath is opt-in
via ``FASTVIDEO_VSA_CUTEDSL=1`` and requires the optional FA4 CuTe
@@ -49,6 +49,26 @@ def _resolve_backend() -> str:
return "triton"
def _expand_mask_and_sizes_128_to_64(
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Expand a [B, H, Qb128, KVb128] map to 64-token Triton tiles."""
expanded_mask = logical_mask_128.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
sizes_i32 = logical_kv_sizes_128.to(torch.int32)
offsets = torch.tensor(
[0, _KV_BLOCK_TRITON],
dtype=torch.int32,
device=sizes_i32.device,
)
expanded_sizes = torch.clamp(
sizes_i32[:, None] - offsets[None, :],
min=0,
max=_KV_BLOCK_TRITON,
).reshape(-1)
return expanded_mask, expanded_sizes
def _expand_mask_and_sizes_256_to_128(
logical_mask_256: torch.Tensor,
logical_kv_sizes_256: torch.Tensor,
@@ -112,6 +132,63 @@ def _triton_via_route_a(
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def _triton_via_route_a_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from .triton_kernels.index import map_to_index as triton_map_to_index
mask_64, sizes_64 = _expand_mask_and_sizes_128_to_64(logical_mask_128, logical_kv_sizes_128)
q2k_idx, q2k_num = triton_map_to_index(mask_64.to(torch.bool))
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def block_sparse_attn_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, H, S, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
return _triton_via_route_a_128(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd
return block_sparse_attn_cute_fwd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_128_bshd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, S, H, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
out_bhsd, aux = _triton_via_route_a_128(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
logical_block_map_128,
logical_variable_block_sizes_128,
)
return out_bhsd.transpose(1, 2).contiguous(), aux
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
return block_sparse_attn_cute_fwd_bshd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_256(
q: torch.Tensor,
k: torch.Tensor,
@@ -1,18 +1,18 @@
"""CuTe-DSL block-sparse attention forward kernel.
"""FA4 CuTe-DSL block-sparse attention adapter.
Thin wrapper around `flash_attn.cute.interface._flash_attn_fwd` that adapts
VSA's `(block_map, variable_block_sizes)` inputs into FA4's
`BlockSparseTensorsTorch` representation and the per-KV-block validity mask.
This module adapts VSA's ``(block_map, variable_block_sizes)`` inputs into
FA4's forward and backward ``BlockSparseTensorsTorch`` representations.
FA4's public ``flash_attn_func`` owns the forward/backward autograd bridge.
Both [B, H, S, D] (BHSD) and [B, S, H, D] (BSHD) entrypoints are provided.
The BSHD variant is preferred from VSA-256 callers to avoid layout
The BSHD variant is preferred from VSA-128/256 callers to avoid layout
round-trips on the hot path.
The FA4 CuTe block-sparse kernel (``flash_attn.cute`` with
``block_sparsity``) is an *optional* dependency: it is imported lazily and
only exercised when the VSA-256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default VSA-256 path is Triton and does
not require it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
only exercised when the VSA-128/256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default path is Triton and does not require
it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
"""
from __future__ import annotations
@@ -22,13 +22,14 @@ from typing import Tuple
import torch
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
_FA4_IMPORT_HINT = ("VSA-128/256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
"provides `flash_attn.cute` with block-sparsity support (plus "
"`nvidia-cutlass-dsl` and `quack-kernels`). This is an optional "
"dependency; the default VSA-256 path is Triton. Install the FA4 CuTe "
"dependency; the default path is Triton. Install the FA4 CuTe "
"build and set FASTVIDEO_VSA_CUTEDSL=1 to enable the CuTe fastpath.")
@functools.lru_cache(maxsize=1)
def _load_fa4_cute():
"""Lazily import the optional FA4 CuTe block-sparse symbols.
@@ -38,14 +39,39 @@ def _load_fa4_cute():
"""
try:
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
from flash_attn.cute.interface import _flash_attn_fwd
from flash_attn.cute.interface import (
_flash_attn_bwd,
_flash_attn_fwd,
flash_attn_func,
)
except ImportError as exc: # pragma: no cover - optional dependency
raise ImportError(_FA4_IMPORT_HINT) from exc
return BlockSparseTensorsTorch, _flash_attn_fwd
return BlockSparseTensorsTorch, flash_attn_func, _flash_attn_fwd, _flash_attn_bwd
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
_M_BLOCK_SIZE_DEFAULT = 128
# FA4's physical Q tile size; KV block size comes from the VSA caller.
_FA4_Q_BLOCK_SIZE = 128
class _SingleQStageLength(int):
"""Keep the real length while selecting FA4's one-stage Q128 path.
On sm_100 FA4 derives ``q_stage`` from ``max_seqlen_q > tile_m``. Its
kernel supports one 128-token Q stage, but the fixed-length public wrapper
does not expose that choice. VSA-128 must select it explicitly; otherwise
adjacent logical Q blocks are merged into a 256-token sparse block.
"""
def __mul__(self, other):
return type(self)(int(self) * int(other))
def __rmul__(self, other):
return type(self)(int(other) * int(self))
def __gt__(self, other):
if int(other) == _FA4_Q_BLOCK_SIZE:
return False
return int(self) > int(other)
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
@@ -64,12 +90,12 @@ def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
return triton_map_to_index(block_map)
def _choose_q_sparse_block_size(q_len: int, m_block_size: int = _M_BLOCK_SIZE_DEFAULT) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > m_block_size.
def _choose_q_sparse_block_size(q_len: int, q_tile_size: int = _FA4_Q_BLOCK_SIZE) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > q_tile_size.
major, _ = torch.cuda.get_device_capability()
if major >= 10 and q_len > m_block_size:
return 2 * m_block_size
return m_block_size
if major >= 10 and q_len > q_tile_size:
return 2 * q_tile_size
return q_tile_size
def _aggregate_q_block_map(
@@ -134,23 +160,35 @@ def _build_vbs_mask_mod(kv_block_size: int):
return _vbs_mask_mod
def _cute_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
def _build_sparse_tensors(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
q_len: int,
q_block_size: int,
kv_block_size: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Internal: FA4 CuTe BSA fwd with BSHD inputs."""
BlockSparseTensorsTorch, _flash_attn_fwd = _load_fa4_cute()
q_sparse_candidate = _choose_q_sparse_block_size(q_bshd.shape[1])
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
need_backward: bool,
force_q_sparse_block_size: int | None = None,
) -> Tuple[object, object | None]:
"""Build the Q-owned forward and KV-owned backward sparse metadata.
``need_backward`` is False on inference-only calls: the backward metadata
is a pair of dense ``[B, H, kv_blocks, q_blocks]`` int32 index tensors that
FA4 keeps alive on its autograd ctx until backward runs, so building it
when nothing requires grad is pure overhead (~80 MiB per call at Wan-14B
720p shape).
"""
BlockSparseTensorsTorch, _, _, _ = _load_fa4_cute()
if force_q_sparse_block_size is None:
q_sparse_candidate = _choose_q_sparse_block_size(q_len)
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
else:
q_sparse_block_size = force_q_sparse_block_size
if q_sparse_block_size < q_block_size or q_sparse_block_size % q_block_size != 0:
raise ValueError("force_q_sparse_block_size must be a positive multiple of q_block_size")
sparse_map = _aggregate_q_block_map(
block_map,
q_sparse_block_size=q_sparse_block_size,
@@ -158,35 +196,166 @@ def _cute_forward(
)
kv_full = (variable_block_sizes == kv_block_size).view(1, 1, 1, -1)
kv_partial = ((variable_block_sizes > 0) & (variable_block_sizes < kv_block_size)).view(1, 1, 1, -1)
full_map = sparse_map & kv_full
mask_map = sparse_map & kv_partial
full_block_idx, full_block_cnt = _map_to_index(full_map)
mask_block_idx, mask_block_cnt = _map_to_index(mask_map)
def from_maps(full_map: torch.Tensor, mask_map: torch.Tensor) -> object:
full_block_idx, full_block_cnt = _map_to_index(full_map.contiguous())
mask_block_idx, mask_block_cnt = _map_to_index(mask_map.contiguous())
return BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
)
sparse_tensors = BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
forward_sparse_tensors = from_maps(
sparse_map & kv_full,
sparse_map & kv_partial,
)
# _flash_attn_fwd returns (out, lse, p, row_max); keep the first two.
out, lse = _flash_attn_fwd(
if not need_backward:
return forward_sparse_tensors, None
# FA4 backward is KV-owned: for each physical KV tile, list the sparse
# query tiles that selected it. Full and partial KV tiles stay separate
# so the token-level validity mask only runs for padded tiles.
backward_sparse_tensors = from_maps(
(sparse_map & kv_full).transpose(2, 3),
(sparse_map & kv_partial).transpose(2, 3),
)
return forward_sparse_tensors, backward_sparse_tensors
def _cute_attention_q128_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
need_backward: bool,
) -> Tuple[torch.Tensor, torch.Tensor, object | None]:
"""Run FA4 with one physical Q stage per logical VSA-128 block."""
_, _, flash_attn_fwd, _ = _load_fa4_cute()
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=_FA4_Q_BLOCK_SIZE,
kv_block_size=_FA4_Q_BLOCK_SIZE,
need_backward=need_backward,
force_q_sparse_block_size=_FA4_Q_BLOCK_SIZE,
)
out, lse = flash_attn_fwd(
q_bshd,
k_bshd,
v_bshd,
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
mask_mod=_build_vbs_mask_mod(kv_block_size),
block_sparse_tensors=sparse_tensors,
tile_mn=(_FA4_Q_BLOCK_SIZE, _FA4_Q_BLOCK_SIZE),
max_seqlen_q=_SingleQStageLength(q_bshd.shape[1]),
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
block_sparse_tensors=forward_sparse_tensors,
aux_tensors=[variable_block_sizes],
causal=False,
return_lse=True,
)[:2]
return out, lse, backward_sparse_tensors
class _CuteAttentionQ128(torch.autograd.Function):
@staticmethod
def forward(ctx, q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes):
out, lse, backward_sparse_tensors = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=True,
)
ctx.save_for_backward(q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes)
ctx.backward_sparse_tensors = backward_sparse_tensors
ctx.mark_non_differentiable(lse)
ctx.set_materialize_grads(False)
return out, lse
@staticmethod
def backward(ctx, grad_out, grad_lse):
del grad_lse
q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes = ctx.saved_tensors
if grad_out is None:
grad_out = torch.zeros_like(out)
_, _, _, flash_attn_bwd = _load_fa4_cute()
dq, dk, dv = flash_attn_bwd(
q_bshd,
k_bshd,
v_bshd,
out,
grad_out.contiguous(),
lse,
softmax_scale=q_bshd.shape[-1]**-0.5,
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=ctx.backward_sparse_tensors,
)
return dq, dk, dv, None, None
def _cute_attention_q128(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
if need_backward:
return _CuteAttentionQ128.apply(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
out, lse, _ = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=False,
)
return out, lse
def _cute_attention(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Run FA4's autograd-enabled block-sparse attention with BSHD inputs."""
_, flash_attn_func, _, _ = _load_fa4_cute()
q_block_size = q_bshd.shape[1] // block_map.shape[2]
kv_block_size = k_bshd.shape[1] // block_map.shape[3]
if q_block_size == kv_block_size == _FA4_Q_BLOCK_SIZE:
return _cute_attention_q128(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=q_block_size,
kv_block_size=kv_block_size,
need_backward=need_backward,
)
return flash_attn_func(
q_bshd,
k_bshd,
v_bshd,
mask_mod=_build_vbs_mask_mod(kv_block_size),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=forward_sparse_tensors,
block_sparse_tensors_bwd=backward_sparse_tensors,
return_lse=True,
)
def block_sparse_attn_cute_fwd(
q: torch.Tensor,
k: torch.Tensor,
@@ -194,34 +363,25 @@ def block_sparse_attn_cute_fwd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""CuTe forward-only block-sparse attention with [B, H, S, D] inputs."""
"""Autograd-enabled CuTe block-sparse attention for [B, H, S, D]."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[2] // block_map.shape[2]
kv_block_size = k.shape[2] // block_map.shape[3]
q_bshd = q.transpose(1, 2).contiguous()
k_bshd = k.transpose(1, 2).contiguous()
v_bshd = v.transpose(1, 2).contiguous()
out_bshd, lse_bshd = _cute_forward(
out_bshd, lse = _cute_attention(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
out = out_bshd.transpose(1, 2).contiguous()
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
# FA4 already returns lse as [B, H, S], matching the Triton path's aux
# contract, so it needs no transpose. Detach before any further op: the
# value is informational and callers never backprop through it.
return out, lse.detach()
def block_sparse_attn_cute_fwd_bshd(
@@ -231,27 +391,16 @@ def block_sparse_attn_cute_fwd_bshd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""CuTe forward-only block-sparse attention with [B, S, H, D] inputs."""
"""Autograd-enabled CuTe block-sparse attention for [B, S, H, D]."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[1] // block_map.shape[2]
kv_block_size = k.shape[1] // block_map.shape[3]
out, lse_bshd = _cute_forward(
out, lse = _cute_attention(
q,
k,
v,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[2], q.shape[1]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
# lse is [B, H, S] regardless of the q/k/v layout; see above.
return out, lse.detach()
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
"""sm_100a (Blackwell) CUDA block-sparse VSA forward.
A third backend behind the same VSA op as the Triton and CuTe-DSL paths. Forward only: it
returns ``(out, lse)`` with ``lse`` in exactly the form ``triton_block_sparse_attn_forward``
writes -- ``max(qk * qk_scale) + log2(l)``, ``[B, H, S]`` fp32 -- so
``block_sparse_attn_backward_triton`` runs against it unchanged.
The extension carries TWO instantiations of the kernel, for 64- and 128-token sparse blocks
(tile volumes 64 and 128 in ``build_vsa_metadata``); the block size is inferred from the
tensors and picks the op. Anything else falls back to Triton via ``is_supported``.
"""
from typing import Tuple
import torch
try:
# The pybind symbols live on fastvideo_kernel_ops, NOT on the _C package that contains it.
# `import fastvideo_kernel._C as _C` resolves to the namespace package, whose __init__ is
# empty, so hasattr() fails on a wheel install and the caller silently falls back with the
# kernel built and present.
from fastvideo_kernel._C import fastvideo_kernel_ops as _C
_FWD_BY_BLOCK = {
64: getattr(_C, "block_sparse_sm100a_fwd", None),
128: getattr(_C, "block_sparse_sm100a_blk128_fwd", None),
}
_HAS_VSA_SM100A = any(_FWD_BY_BLOCK.values())
except ImportError: # pragma: no cover - extension not built
_C = None
_FWD_BY_BLOCK = {}
_HAS_VSA_SM100A = False
_SM100 = (10, 0)
HEAD_DIM = 128
# Must match the -DVSA_BHSD the extension was compiled with (see CMakeLists).
BHSD = True
def _block_size(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> int:
num_blocks = variable_block_sizes.numel()
seqlen = q.shape[2] if BHSD else q.shape[1]
return 0 if num_blocks == 0 or seqlen % num_blocks else seqlen // num_blocks
def is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
"""True iff this build can run these tensors; otherwise the caller uses Triton.
Static facts only -- shapes, dtypes, arch, layout. Deliberately NO reads of tensor
contents: the previous ``int(variable_block_sizes.min())`` was a GPU->CPU sync on every
call, and the kernel no longer needs it (see below). This predicate must stay cheap
enough to sit on a per-layer dispatch path.
What the kernel accepts (and is tested to handle):
* q/k/v: contiguous 4-D bf16 CUDA tensors on an sm_100 device, head_dim 128, laid out
as compiled (BHSD here); seqlen == num_blocks * block with an EVEN num_blocks (a CTA
owns an adjacent pair of query blocks) and a 64- or 128-token build present.
* q2k_num: any per-row counts in [0, max_kv], NON-uniform across rows included. Rows
with count 0 produce exactly-zero output rows (and a finite LSE sentinel) rather
than attending anywhere -- so no ``.min()`` floor is required of the caller.
* q2k_idx: rows only need valid entries (in [0, num_blocks)) BELOW that row's count;
padding past the count (e.g. map_to_index's -1 fill) is never dereferenced. max_kv
(= q2k_idx.shape[-1]) must be >= 1, which the host launcher re-checks.
* variable_block_sizes: per-KV-block valid-token counts in [0, block]; keys at or past
a block's count are masked. Integer metadata is converted to int32/contiguous by
``block_sparse_attn_sm100a`` itself, so int64 inputs merely cost a cast.
"""
if not _HAS_VSA_SM100A or not q.is_cuda:
return False
if torch.cuda.get_device_capability(q.device) != _SM100:
return False
if q.dtype != torch.bfloat16 or q.dim() != 4 or q.shape[-1] != HEAD_DIM:
return False
if not q.is_contiguous():
return False
if _FWD_BY_BLOCK.get(_block_size(q, variable_block_sizes)) is None:
return False
# A CTA owns an adjacent pair of query blocks.
if variable_block_sizes.numel() % 2 != 0:
return False
# Metadata must be integer-typed so the wrapper's int32 conversion is value-preserving.
if not variable_block_sizes.is_cuda or variable_block_sizes.dtype not in (torch.int32, torch.int64):
return False
return True
def block_sparse_attn_sm100a(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
need_lse: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass. Returns ``(out, lse)``; ``out`` has q's layout."""
fwd = _FWD_BY_BLOCK[_block_size(q, variable_block_sizes)]
idx = q2k_idx.to(torch.int32).contiguous()
num = q2k_num.to(torch.int32).contiguous()
vbs = variable_block_sizes.to(torch.int32).contiguous()
sm_scale = 1.0 / (q.shape[-1]**0.5)
res = fwd(q.contiguous(), k.contiguous(), v.contiguous(), None,
idx, num, vbs, sm_scale, need_lse)
return (res[0], res[1]) if need_lse else (res[0], None)
+24 -22
View File
@@ -2,6 +2,8 @@ import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn_256 import (
block_sparse_attn_128,
block_sparse_attn_128_bshd,
block_sparse_attn_256,
block_sparse_attn_256_bshd,
)
@@ -74,12 +76,13 @@ def video_sparse_attn(
Dispatches the sparse branch by ``block_elements = prod(block_size)``:
- 64 -> existing TK/Triton path (see ``block_sparse_attn_from_indices``).
- 128 -> Triton fallback or CuTe FA4 block-sparse attention.
- 256 -> CuTe FA4 block-sparse attention (see ``block_sparse_attn_256``).
Backend overrides:
- ``FASTVIDEO_VSA_TRITON=1`` forces Triton in either path.
- ``FASTVIDEO_VSA_TK=1`` prefers the sm_90 TK kernel in the 64-block path.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 256-block path.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 128/256-block paths.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
@@ -119,8 +122,9 @@ def video_sparse_attn(
# Sparse branch (fused Triton topk mask)
mask = fused_topk_mask(scores, topk)
if block_elements == 256:
out_s = block_sparse_attn_256(q, k, v, mask, variable_block_sizes)[0]
if block_elements in (128, 256):
attention = block_sparse_attn_128 if block_elements == 128 else block_sparse_attn_256
out_s = attention(q, k, v, mask, variable_block_sizes)[0]
else:
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
@@ -142,14 +146,14 @@ def video_sparse_attn_bshd(
"""VSA entrypoint for [B, S, H, D] tensors.
Avoids the BHSD<->BSHD round-trip that ``video_sparse_attn`` performs on
the CuTe 256-block path; the 64-block path still expects BHSD and is not
the CuTe 128/256-block paths; the 64-block path still expects BHSD and is not
supported here.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
if block_elements != 256:
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=256 "
if block_elements not in (128, 256):
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=128 or 256 "
f"(got {block_elements}); use video_sparse_attn for the 64-block path.")
batch, q_seq_len, heads, dim = q.shape
@@ -171,19 +175,15 @@ def video_sparse_attn_bshd(
raise ValueError(f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
f"got {q_variable_block_sizes.numel()}")
# Compression branch (BSHD-native: mean over the 256-token axis).
token_idx = torch.arange(block_elements, device=q.device, dtype=torch.int32)
q_token_valid = (token_idx.view(1, -1) < q_variable_block_sizes.view(-1,
1)).view(1, q_num_blocks, block_elements, 1, 1)
kv_token_valid = (token_idx.view(1, -1) < variable_block_sizes.view(-1,
1)).view(1, kv_num_blocks, block_elements, 1, 1)
# Compression branch (BSHD-native: match fused_block_mean's semantics).
# Padding values are expected to be zero; gradients are broadcast across
# the full padded block, just like the BHSD fused common path.
q_c = q.view(batch, q_num_blocks, block_elements, heads, dim)
k_c = k.view(batch, kv_num_blocks, block_elements, heads, dim)
v_c = v.view(batch, kv_num_blocks, block_elements, heads, dim)
q_c = ((q_c.float() * q_token_valid).sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = ((k_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = ((v_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_c = (q_c.float().sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = (k_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = (v_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_ch = q_c.permute(0, 2, 1, 3).contiguous()
k_ch = k_c.permute(0, 2, 1, 3).contiguous()
v_ch = v_c.permute(0, 2, 1, 3).contiguous()
@@ -195,13 +195,15 @@ def video_sparse_attn_bshd(
# Sparse branch (fused Triton topk mask + CuTe BSHD).
mask = fused_topk_mask(scores, topk)
out_s, _ = block_sparse_attn_256_bshd(q, k, v, mask, variable_block_sizes)
attention = block_sparse_attn_128_bshd if block_elements == 128 else block_sparse_attn_256_bshd
out_s, _ = attention(q, k, v, mask, variable_block_sizes)
out = out_s
out_view = out.view(batch, q_num_blocks, block_elements, heads, dim)
# Out-of-place: ``out_s`` is the tensor FA4's autograd node saved for its
# backward, so mutating it in place invalidates the graph.
out_view = out_s.view(batch, q_num_blocks, block_elements, heads, dim)
if compress_attn_weight is not None:
gate_view = compress_attn_weight.view(batch, q_num_blocks, block_elements, heads, dim)
out_view.add_(out_c_blk.unsqueeze(2) * gate_view)
out = out_view + out_c_blk.unsqueeze(2) * gate_view
else:
out_view.add_(out_c_blk.unsqueeze(2))
return out
out = out_view + out_c_blk.unsqueeze(2)
return out.view(batch, q_seq_len, heads, dim)
@@ -237,7 +237,12 @@ def _attn_bwd_dkdv(
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
qkT = tl.dot(k, qT)
# Recompute logits exactly as the forward does: raw bf16 operands into
# the dot, fp32 scale after accumulation. A bf16 pre-scaled K perturbs
# the recomputed logits relative to the saved M by an error
# proportional to |logit|, which exp2 amplifies into arbitrarily wrong
# probabilities at large activations.
qkT = tl.dot(k, qT) * (sm_scale * 1.4426950408889634)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
@@ -268,6 +273,7 @@ def _attn_bwd_dq(
do,
m,
D,
sm_scale,
# shared by Q/K/V/DO.
q2k_index,
q2k_num,
@@ -315,7 +321,7 @@ def _attn_bwd_dq(
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
p = tl.math.exp2(qk - m)
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
mask = offs_in_block < block_size
@@ -324,8 +330,7 @@ def _attn_bwd_dq(
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@@ -453,6 +458,7 @@ def _attn_bwd(
do,
m,
D, #
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
@@ -470,7 +476,7 @@ def _attn_bwd(
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= LN2
dq *= sm_scale
tl.store(dq_ptrs, dq)
@@ -591,6 +597,7 @@ def _attn_bwd_dq_kernel(
Q,
K,
V,
sm_scale,
DO, #
DQ,
M,
@@ -663,6 +670,7 @@ def _attn_bwd_dq_kernel(
do,
m,
D,
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
@@ -680,7 +688,7 @@ def _attn_bwd_dq_kernel(
)
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq_acc *= LN2
dq_acc *= sm_scale
tl.store(dq_ptrs, dq_acc)
@@ -748,9 +756,11 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
dv = torch.empty_like(v)
BATCH, N_HEAD = q.shape[:2]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
# K stays raw: the backward kernels apply sm_scale in fp32 after the dot,
# matching the forward's rounding exactly. (A bf16 pre-scaled K perturbs
# the recomputed logits vs the saved M; exp2 turns that into unboundedly
# wrong probabilities at large activations.)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert Tq % PRE_BLOCK == 0
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
@@ -813,6 +823,7 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
q,
arg_k,
v,
sm_scale,
do,
dq,
M,
@@ -14,7 +14,9 @@ import math
import torch
VSA_TILE_SIZE = (4, 4, 4)
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 256)
# 128 is served by the sm_100a CUDA backend (block_sparse_attn_sm100a); 64 and 256 by
# Triton and the CuTe-DSL path. A volume here only needs a backend that accepts it.
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 128, 256)
def _canonicalize_device(device: torch.device | str) -> torch.device:
@@ -0,0 +1,146 @@
"""VSA-128 CuTe/Triton forward and backward parity on Blackwell."""
from __future__ import annotations
import math
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 128
_BLOCK_SIZE_3D = (2, 8, 8)
def _select_backend(monkeypatch, backend: str) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _dense_sparse_reference(q, k, v, block_map, variable_block_sizes):
token_mask = block_map.repeat_interleave(_BLOCK, dim=2).repeat_interleave(_BLOCK, dim=3)
kv_valid = torch.arange(_BLOCK, device=k.device) < variable_block_sizes[:, None]
token_mask = token_mask & kv_valid.reshape(1, 1, 1, -1)
logits = torch.matmul(q.float(), k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
probabilities = torch.softmax(logits.masked_fill(~token_mask, float("-inf")), dim=-1)
return torch.matmul(probabilities, v.float()).to(q.dtype)
def _check(tag: str, expected: torch.Tensor, actual: torch.Tensor, avg_tol: float, rel_tol: float) -> None:
assert torch.isfinite(actual).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(expected, actual)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < avg_tol
assert max_rel < rel_tol
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
def test_vsa128_explicit_routes_forward_backward(backend: str, monkeypatch) -> None:
"""Adjacent Q128 blocks must keep independent routes instead of merging."""
_select_backend(monkeypatch, backend)
torch.manual_seed(53)
shape = (1, 1, 3 * _BLOCK, 128)
base = [torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3)]
grad_output = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
variable_block_sizes = torch.tensor([128, 91, 37], device="cuda", dtype=torch.int32)
block_map = torch.eye(3, device="cuda", dtype=torch.bool).view(1, 1, 3, 3)
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
actual, _ = block_sparse_attn_128(*actual_inputs, block_map, variable_block_sizes)
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
expected = _dense_sparse_reference(*reference_inputs, block_map, variable_block_sizes)
(expected * grad_output).sum().backward()
print(f"[vsa128-explicit-{backend}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_inputs, strict=True):
_check(name, reference.grad, candidate.grad, 2e-2, 0.5)
def _zero_kv_tail(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
valid = torch.arange(_BLOCK, device=x.device) < variable_block_sizes[:, None]
valid = valid.view(1, 1, -1, _BLOCK, 1).expand_as(x.view(1, x.shape[1], -1, _BLOCK, x.shape[-1]))
return x * valid.reshape_as(x).to(x.dtype)
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa128_wrapper_forward_backward(backend: str, layout: str, monkeypatch) -> None:
_select_backend(monkeypatch, backend)
torch.manual_seed(59)
batch, heads, dim = 1, 2, 128
q_blocks, kv_blocks, topk = 3, 4, 2
q_shape = (batch, heads, q_blocks * _BLOCK, dim)
kv_shape = (batch, heads, kv_blocks * _BLOCK, dim)
q_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
kv_sizes = torch.tensor([128, 91, 37, 128], device="cuda", dtype=torch.int32)
q_sizes = torch.full((q_blocks, ), _BLOCK, device="cuda", dtype=torch.int32)
k_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
v_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
gate_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16) * 0.1
grad_output = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
if layout == "bhsd":
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
actual_gate = gate_base.detach().clone().requires_grad_()
actual = video_sparse_attn(
*actual_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=actual_gate,
)
actual_grads = actual_inputs
else:
bshd_inputs = [tensor.transpose(1, 2).contiguous().detach().requires_grad_()
for tensor in (q_base, k_base, v_base)]
bshd_gate = gate_base.transpose(1, 2).contiguous().detach().requires_grad_()
actual = video_sparse_attn_bshd(
*bshd_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=bshd_gate,
).transpose(1, 2)
actual_grads = bshd_inputs
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
reference_gate = gate_base.detach().clone().requires_grad_()
expected = _torch_vsa256_reference(
*reference_inputs,
q_sizes,
kv_sizes,
topk,
compress_attn_weight=reference_gate,
)
(expected * grad_output).sum().backward()
print(f"[vsa128-wrapper-{backend}-{layout}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_grads, strict=True):
candidate_grad = candidate.grad if layout == "bhsd" else candidate.grad.transpose(1, 2)
_check(name, reference.grad, candidate_grad, 2e-2, 0.5)
actual_gate_grad = actual_gate.grad if layout == "bhsd" else bshd_gate.grad.transpose(1, 2)
_check("dgate", reference_gate.grad, actual_gate_grad, 1e-3, 0.2)
@@ -0,0 +1,224 @@
"""VSA-256 FA4 CuTe forward/backward parity for BHSD and BSHD APIs.
Covers the shapes the CuTe backward actually sees in production: the gated
compression branch (`compress_attn_weight`), partially filled Q tiles,
and q_len != kv_len. Also pins the inference fast path, which must skip the
KV-owned backward metadata without changing the forward result.
"""
from __future__ import annotations
from typing import Tuple
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 256
_BLOCK_SIZE_3D = (4, 8, 8) # prod == 256
# Measured on GB200 (sm_100) with bf16 inputs: grads land around 1e-4 avg_abs
# and <=0.11 max_rel across every case below, so these leave ~10x headroom
# without being loose enough to hide a real regression.
_OUT_TOL = (1e-3, 0.2)
_GRAD_TOL = (1e-3, 0.25)
@pytest.fixture(autouse=True)
def _require_cute_backend(monkeypatch):
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
def _zero_pad_tail(x: torch.Tensor, var: torch.Tensor) -> torch.Tensor:
"""Zero the padded tail of every 256-token tile of a [B, H, S, D] tensor.
VSA callers scatter into a zeroed tile buffer, so padded slots are zero;
both the kernel and the reference rely on that.
"""
bsz, heads, _, dim = x.shape
blocks = var.numel()
token_idx = torch.arange(_BLOCK, device=x.device, dtype=torch.int32)
valid = (token_idx.view(1, -1) < var.view(-1, 1)).view(1, 1, blocks, _BLOCK, 1)
valid = valid.expand(bsz, heads, blocks, _BLOCK, dim).reshape_as(x)
return x * valid.to(x.dtype)
def _make_inputs(
q_blocks: int,
kv_blocks: int,
kv_var: torch.Tensor,
q_var: torch.Tensor,
heads: int = 2,
dim: int = 128,
seed: int = 42,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch.manual_seed(seed)
device = torch.device("cuda")
dtype = torch.bfloat16
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=dtype)
k = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
v = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
return _zero_pad_tail(q, q_var), _zero_pad_tail(k, kv_var), _zero_pad_tail(v, kv_var), grad_out
def _check(tag: str, ref: torch.Tensor, got: torch.Tensor, tol: Tuple[float, float]) -> None:
assert torch.isfinite(got).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(ref, got)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < tol[0], f"{tag}: avg_abs {avg_abs:.3e} >= {tol[0]:.3e}"
assert max_rel < tol[1], f"{tag}: max_rel {max_rel:.3e} >= {tol[1]:.3e}"
def _run_bhsd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = video_sparse_attn(qg, kg, vg, kv_var, q_var, topk, block_size=_BLOCK_SIZE_3D, compress_attn_weight=gate)
return out, (qg, kg, vg)
def _run_bshd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.transpose(1, 2).contiguous().requires_grad_(True) for t in (q, k, v))
gate_bshd = None if gate is None else gate.transpose(1, 2).contiguous()
out = video_sparse_attn_bshd(qg,
kg,
vg,
kv_var,
q_var,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=gate_bshd)
return out.transpose(1, 2), (qg, kg, vg)
def _reference(q, k, v, q_var, kv_var, topk, gate=None):
qr, kr, vr = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = _torch_vsa256_reference(qr, kr, vr, q_var, kv_var, topk, compress_attn_weight=gate)
return out, (qr, kr, vr)
def _compare(tag, layout, q, k, v, kv_var, q_var, topk, grad_out, gate=None):
runner = _run_bhsd if layout == "bhsd" else _run_bshd
out, (qg, kg, vg) = runner(q, k, v, kv_var, q_var, topk, gate=gate)
(out * grad_out).sum().backward()
grads = [g.grad if g.grad.dim() == 4 and layout == "bhsd" else g.grad for g in (qg, kg, vg)]
if layout == "bshd":
grads = [g.transpose(1, 2) for g in grads]
out_ref, refs = _reference(q, k, v, q_var, kv_var, topk, gate=gate)
(out_ref * grad_out).sum().backward()
print(f"[{tag}-{layout}]")
_check("out", out_ref, out, _OUT_TOL)
for name, ref, got in zip(("dq", "dk", "dv"), refs, grads):
_check(name, ref.grad, got, _GRAD_TOL)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_forward_backward_vs_torch_ref(layout: str) -> None:
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var)
_compare("vsa256-cute", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_with_compress_gate(layout: str) -> None:
"""The gated compression branch is what Wan and MiniMax-H3 actually run.
It is also the branch that composes the sparse output with the compression
output, so it is the one that breaks if that composition mutates FA4's
saved output in place.
"""
kv_var = torch.tensor([256, 200, 256, 91], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var, seed=7)
gate = torch.randn_like(q) * 0.1
_compare("vsa256-cute-gated", layout, q, k, v, kv_var, q_var, 2, grad_out, gate=gate)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_partial_q_blocks(layout: str) -> None:
"""Q tiles that are not full: only the compression divisor depends on it,
but it is the one axis the existing coverage held constant."""
kv_var = torch.tensor([256, 128, 256], dtype=torch.int32, device="cuda")
q_var = torch.tensor([256, 61, 199], dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 3, kv_var, q_var, seed=11)
_compare("vsa256-cute-partial-q", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_cross_q_kv(layout: str) -> None:
"""q_len != kv_len: forward has coverage, backward did not."""
kv_var = torch.tensor([256, 143, 256, 256, 88], dtype=torch.int32, device="cuda")
q_var = torch.full((2, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(2, 5, kv_var, q_var, seed=13)
_compare("vsa256-cute-cross", layout, q, k, v, kv_var, q_var, 3, grad_out)
@pytest.mark.cuda
def test_vsa256_cute_inference_matches_training_forward() -> None:
"""The KV-owned backward metadata is only built when something requires
grad. Skipping it must not perturb the forward result."""
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, _ = _make_inputs(3, 4, kv_var, q_var, seed=5)
with torch.no_grad():
out_infer = video_sparse_attn_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
kv_var,
q_var,
2,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=None,
)
out_train, _ = _run_bshd(q, k, v, kv_var, q_var, 2)
torch.testing.assert_close(out_infer, out_train.transpose(1, 2).detach(), rtol=0, atol=0)
@pytest.mark.cuda
def test_vsa256_cute_lse_is_bhs() -> None:
"""The aux return is [B, H, S] on both entrypoints, matching the Triton
path's contract."""
from fastvideo_kernel.block_sparse_attn_256 import (block_sparse_attn_256, block_sparse_attn_256_bshd)
device = torch.device("cuda")
heads, dim, q_blocks, kv_blocks = 2, 128, 3, 4
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=torch.bfloat16)
k = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
v = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
vbs = torch.full((kv_blocks, ), _BLOCK, dtype=torch.int32, device=device)
mask = torch.zeros(1, heads, q_blocks, kv_blocks, dtype=torch.bool, device=device)
mask[..., :2] = True
_, lse_bhsd = block_sparse_attn_256(q, k, v, mask, vbs)
assert lse_bhsd.shape == (1, heads, sq), lse_bhsd.shape
_, lse_bshd = block_sparse_attn_256_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
mask,
vbs,
)
assert lse_bshd.shape == (1, heads, sq), lse_bshd.shape
@@ -22,6 +22,7 @@ def _torch_vsa256_reference(
q_var: torch.Tensor,
kv_var: torch.Tensor,
topk_logical: int,
compress_attn_weight: torch.Tensor | None = None,
) -> torch.Tensor:
bsz, heads, _sq, dim = q.shape
q_blocks = q_var.numel()
@@ -55,6 +56,8 @@ def _torch_vsa256_reference(
logits = logits.masked_fill(~token_mask, float("-inf"))
prob = torch.softmax(logits, dim=-1)
out_s = torch.matmul(prob, vf).to(q.dtype)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
return out_c + out_s
@@ -0,0 +1,104 @@
"""Regression: Triton block-sparse backward gradient parity at realistic activation scale.
The backward used to fold ``sm_scale / ln(2)`` into K in bf16 before the
exp2-based logit recompute. The bf16 rounding error on the pre-scaled K grows
proportionally to |logit| and exp2 amplifies it into exponentially wrong
probabilities, so dQ/dK/dV were correct at unit scale (every pre-existing test)
but off by orders of magnitude at real activation magnitudes.
This test sweeps the input scale and checks the Triton kernel's gradients
against an fp32 masked-dense SDPA reference. The unit-scale case is the
control (it passed even with the broken kernel); the large-scale cases are
the regression.
"""
import pytest
import torch
from fastvideo_kernel.block_sparse_attn import _map_to_index, block_sparse_attn_triton
from .utils import generate_block_sparse_mask_for_function
BLOCK = 64
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Same convention as test_vsa_varlen.py: every tensor here comes from the
global torch RNG and the checks use tight thresholds, so an unseeded run
would shift inputs whenever an earlier test file draws a different number
of randoms.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _dense_reference(q, k, v, block_mask):
"""fp32 masked-dense SDPA over the token-expanded block mask.
q/k/v: [B, H, S, D]; block_mask: [B, H, S // BLOCK, S // BLOCK] bool.
"""
qf, kf, vf = q.float(), k.float(), v.float()
token_mask = block_mask.repeat_interleave(BLOCK, dim=-2).repeat_interleave(BLOCK, dim=-1)
logits = torch.matmul(qf, kf.transpose(-2, -1)) * (q.shape[-1]**-0.5)
logits = logits.masked_fill(~token_mask, float("-inf"))
return torch.matmul(logits.softmax(dim=-1), vf)
@pytest.mark.cuda
@pytest.mark.parametrize("scale", [1.0, 4.0, 16.0])
def test_triton_backward_grad_parity_across_input_scales(scale: float) -> None:
"""Kernel dQ/dK/dV must stay within a few percent of the fp32 reference
regardless of input magnitude.
With the bf16 K pre-scaling bug, scale<=4.0 passes at this geometry while
scale=16.0 fails (measured on GB200: dq relative L2 error 5.9e-1 vs 6.9e-3
fixed); at larger geometries and real activation magnitudes the broken
kernel is off by orders of magnitude. The passing unit-scale case is
exactly how the bug survived the original test suite.
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
device = torch.device("cuda")
dtype = torch.bfloat16
batch, heads, dim = 1, 4, 128
num_blocks = 8
seq = num_blocks * BLOCK
q = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
k = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
v = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
block_mask = generate_block_sparse_mask_for_function(heads, num_blocks, num_blocks, k=3,
device=device).unsqueeze(0)
q2k_idx, q2k_num = _map_to_index(block_mask)
variable_block_sizes = torch.full((num_blocks, ), BLOCK, dtype=torch.int32, device=device)
q_ker, k_ker, v_ker = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out_ker, _ = block_sparse_attn_triton(q_ker, k_ker, v_ker, q2k_idx, q2k_num, variable_block_sizes)
(out_ker.float() * grad_out.float()).sum().backward()
q_ref, k_ref, v_ref = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out_ref = _dense_reference(q_ref, k_ref, v_ref, block_mask)
(out_ref * grad_out.float()).sum().backward()
# Forward is exact at any scale; this pins the harness itself.
fwd_rel = ((out_ker.float() - out_ref).norm() / out_ref.norm()).item()
assert fwd_rel < 2e-2, f"scale={scale}: forward rel err {fwd_rel:.3e}"
for name, g_ker, g_ref in (
("dq", q_ker.grad, q_ref.grad),
("dk", k_ker.grad, k_ref.grad),
("dv", v_ker.grad, v_ref.grad),
):
assert torch.isfinite(g_ker).all().item(), f"scale={scale}: non-finite {name}"
ref_norm = g_ref.float().norm()
rel = ((g_ker.float() - g_ref.float()).norm() / ref_norm.clamp_min(1e-12)).item()
ratio = (g_ker.float().norm() / ref_norm.clamp_min(1e-12)).item()
print(f"scale={scale} {name}: rel_l2={rel:.4e} norm_ratio={ratio:.4f}")
assert rel < 5e-2, f"scale={scale}: {name} rel l2 err {rel:.3e} >= 5e-2"
assert 0.98 < ratio < 1.02, f"scale={scale}: {name} grad-norm ratio {ratio:.4f}"
+14
View File
@@ -23,6 +23,20 @@ from fastvideo_kernel.block_sparse_attn import (
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Every tensor and every variable block size here comes from the global
torch RNG, and the gradient checks use a tight max_rel threshold. Without
a seed the inputs shift whenever an earlier test file draws a different
number of randoms, which surfaces as an unrelated-looking failure in
whichever case happens to land on unlucky data.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _reference_per_sequence(
q_list,
k_list,
+6 -1
View File
@@ -19,7 +19,12 @@ from fastvideo.attention.backends.abstract import (
from fastvideo.logger import init_logger
logger = init_logger(__name__)
logger.info("Using FlashAttention-%s backend", fa_version)
# Every worker records the loaded FlashAttention implementation so a
# distributed profiling log contains one backend receipt per rank.
logger.info("Worker %s Using FlashAttention-%s backend",
os.environ.get("RANK", "0"),
fa_version,
local_main_process_only=False)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
@@ -5,13 +5,17 @@ H3 runs one joint bidirectional attention over
``[text | condition keyframes | audio | generated video]``, so this
backend differs from the Wan-tuned ``video_sparse_attn``:
- Tiles are ``[segment-pure prefix chunks] + [3D (4,8,8) video tiles]``;
prefix tiles never straddle segment boundaries.
- Tiles are ``[segment-pure prefix chunks] + [3D video tiles]``; prefix
tiles never straddle segment boundaries. The tile size is selectable at
metadata build time: 256 tokens ``(4,8,8)`` (default) or 64 tokens
``(4,4,4)`` (see ``VSA_H3_TILE_SHAPES``).
- Selection is pure Python on pooled tile scores; the block-sparse kernel
consumes an explicit bool mask, so no kernel changes are needed.
- The compression branch is gated by ``to_gate_compress``, which the H3
checkpoint does not carry: the loader zero-initializes it, so untrained
inference is exactly pure sparse and finetuning can learn the gate.
- The compression branch is gated by ``to_gate_compress``, which the base
H3 checkpoint does not carry: the loader zero-initializes it, so
untrained inference is exactly pure sparse and finetuning can learn the
gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship
trained gates, which load and activate the branch.
- Non-video *queries* are always dense. Non-video *keys* are either
always-selected for every query ("exempt", default) or compete in
top-k under a FLOP-matched budget ("compete") — the ablation axis,
@@ -20,22 +24,46 @@ backend differs from the Wan-tuned ``video_sparse_attn``:
(``vsa_dense_first_n_steps``, ``vsa_dense_layers``) let mixed schedules
run the diffuse steps/layers dense while pushing the rest harder.
Targets sm10.x through the FA4 CuTe 256-tile path
At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path
(``FASTVIDEO_VSA_CUTEDSL=1``); the Triton 256→64 expansion is the
fallback and keeps identical mask semantics.
fallback and keeps identical mask semantics. At tile 64 the block map is
already at the kernels' native 64-token granularity, so both forward and
backward run the Triton block-sparse kernels directly (no expansion,
``FASTVIDEO_VSA_CUTEDSL`` does not apply). A third, opt-in route exists
for the tile-64 FORWARD only: ``FASTVIDEO_VSA_SM100A=1`` sends no-grad
forwards through the sm_100a CUDA block-sparse kernel
(``fastvideo_kernel.block_sparse_attn_sm100a``, upstream PR #1719 plus
our per-q-tile ``q2k_num`` fix) when the extension is built, the device
is sm_100, and the geometry qualifies; grad-tracking forwards and every
backward stay on Triton unchanged. If the env is set but a precondition
fails, the route logs one warning and falls back.
"""
import functools
import math
import os
from dataclasses import dataclass
from typing import Any
import torch
try:
from fastvideo_kernel.block_sparse_attn import block_sparse_attn as block_sparse_attn_64_bhsd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
from fastvideo_kernel.triton_kernels.index import map_to_index
except ImportError:
block_sparse_attn_64_bhsd = None
block_sparse_attn_256_bshd = None
map_to_index = None
try:
# Optional: only present in fastvideo_kernel builds that carry the sm_100a
# CUDA block-sparse forward (upstream PR #1719). The module itself imports
# fine without the compiled symbols (`_HAS_VSA_SM100A` is then False and
# `is_supported` says no), so this only guards *module* availability.
from fastvideo_kernel import block_sparse_attn_sm100a as _sm100a
except ImportError:
_sm100a = None
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder, layer_idx_from_prefix)
@@ -43,51 +71,115 @@ from fastvideo.attention.backends.video_sparse_attn import (compute_topk, constr
get_non_pad_index, get_tile_partition_indices,
scatter_into_tile_buf)
from fastvideo.attention.backends.video_sparse_attn_h3_probe import probe_enabled, record_probe
from fastvideo.logger import init_logger
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
logger = init_logger(__name__)
# Opt-in switch for the sm_100a CUDA forward on the tile-64 no-grad path.
VSA_SM100A_ENV = "FASTVIDEO_VSA_SM100A"
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x (default)
_TILE_ELEMS = math.prod(VSA_H3_TILE_SIZE)
# Selectable tile geometries, keyed by element count (= the build-time
# ``tile_size``). 64 runs the native 64-token Triton block-sparse kernels for
# forward AND backward — the block map is already at kernel granularity, so no
# 256->64 mask expansion is involved and FASTVIDEO_VSA_CUTEDSL does not apply.
VSA_H3_TILE_SHAPES: dict[int, tuple[int, int, int]] = {
_TILE_ELEMS: VSA_H3_TILE_SIZE,
64: (4, 4, 4),
}
def token_tile_and_valid(variable_block_sizes: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
"""Per padded-token tile id and pad-validity mask.
The single encoding of the padding contract, shared by the probe and the
test oracle so they cannot drift from the backend's tile geometry.
``tile_elems`` must match the metadata the sizes came from
(``MiniMaxH3VSAMetadata.tile_elems``).
"""
device = variable_block_sizes.device
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(_TILE_ELEMS)
token_valid = (torch.arange(_TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
return token_tile, token_valid
def _validate_h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
variable_block_sizes: torch.Tensor,
untile_combined_index: torch.Tensor,
tile_elems: int = _TILE_ELEMS,
) -> None:
"""Fail synchronously on out-of-bounds tile geometry.
Invariants the block-sparse kernel trusts without checking:
every tile's valid size is in (0, tile_elems]; the sizes sum to the
packed sequence length; and ``untile_combined_index`` maps each packed
row to exactly one non-pad slot of the padded tile buffer. A violation
would surface only as an async device fault at some later kernel or
collective (e.g. an FSDP all-gather), which is unattributable — so raise
here, once per cached geometry, with the numbers in hand.
"""
total = sum(prefix_segments) + math.prod(dit_seq_shape)
n_pad = variable_block_sizes.numel() * tile_elems
sizes_min = int(variable_block_sizes.min())
sizes_max = int(variable_block_sizes.max())
sizes_sum = int(variable_block_sizes.sum())
if sizes_min < 1 or sizes_max > tile_elems or sizes_sum != total:
raise ValueError(f"VSA-H3 tile sizes out of bounds for prefix={prefix_segments}, video={dit_seq_shape}, "
f"tile_elems={tile_elems}: min={sizes_min}, max={sizes_max}, sum={sizes_sum}, "
f"expected sum={total}.")
if untile_combined_index.numel() != total:
raise ValueError(f"VSA-H3 untile index has {untile_combined_index.numel()} entries for a packed "
f"sequence of {total} rows (prefix={prefix_segments}, video={dit_seq_shape}).")
idx_min = int(untile_combined_index.min())
idx_max = int(untile_combined_index.max())
if idx_min < 0 or idx_max >= n_pad:
# Range first: the pad-slot gather below would itself index out of
# bounds (the very async fault this guard exists to preempt).
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: range "
f"[{idx_min}, {idx_max}] vs padded length {n_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
in_tile_offset = untile_combined_index % tile_elems
maps_into_pad = bool((in_tile_offset >= variable_block_sizes[untile_combined_index // tile_elems]).any())
if maps_into_pad or int(torch.unique(untile_combined_index).numel()) != total:
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: "
f"pad-slot hit={maps_into_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
@functools.lru_cache(maxsize=10)
def _h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
device: torch.device,
tile_shape: tuple[int, int, int] = VSA_H3_TILE_SIZE,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
"""Tile the packed sequence: segment-pure prefix chunks, then video tiles.
Returns (tile_partition_indices, variable_block_sizes,
untile_combined_index, num_prefix_tiles, num_video_tiles).
"""
tile_elems = math.prod(tile_shape)
prefix_len = sum(prefix_segments)
prefix_sizes: list[int] = []
for segment in prefix_segments:
full, rem = divmod(segment, _TILE_ELEMS)
prefix_sizes.extend([_TILE_ELEMS] * full)
full, rem = divmod(segment, tile_elems)
prefix_sizes.extend([tile_elems] * full)
if rem:
prefix_sizes.append(rem)
num_prefix_tiles = len(prefix_sizes)
ts_t, ts_h, ts_w = VSA_H3_TILE_SIZE
ts_t, ts_h, ts_w = tile_shape
t, h, w = dit_seq_shape
num_tiles = (math.ceil(t / ts_t), math.ceil(h / ts_h), math.ceil(w / ts_w))
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, VSA_H3_TILE_SIZE)
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
num_video_tiles = int(video_sizes.numel())
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
tile_partition_indices = torch.cat([
torch.arange(prefix_len, device=device, dtype=torch.long),
video_indices,
@@ -100,9 +192,11 @@ def _h3_tile_geometry(
# get_non_pad_index is lru-cached on tensor identity; variable_block_sizes
# is itself cached by this function, so the identity stays stable.
non_pad_index = get_non_pad_index(variable_block_sizes, _TILE_ELEMS)
non_pad_index = get_non_pad_index(variable_block_sizes, tile_elems)
untile_combined_index = non_pad_index[torch.argsort(tile_partition_indices)]
# One-time (lru-cached) synchronous bounds check; see _validate_h3_tile_geometry.
_validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
return (tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles)
@@ -139,6 +233,9 @@ class MiniMaxH3VSAMetadata(AttentionMetadata):
exempt: bool
variable_block_sizes: torch.Tensor
untile_combined_index: torch.Tensor
# tokens per tile (256 or 64); selects the tile geometry AND the kernel
# route in forward() (256 -> VSA-256 CuTe/Triton, 64 -> native Triton)
tile_elems: int = _TILE_ELEMS
# layers forced dense regardless of sparsity (probe-guided opt-outs)
dense_layers: tuple[int, ...] = ()
# Single-slot holder for the padded tile buffer, owned by the BUILDER so
@@ -158,24 +255,28 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
prefix_segments: tuple[int, ...],
device: torch.device,
exempt: bool = True,
dense_layers: tuple[int, ...] = (),
**kwargs: dict[str, Any],
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
prefix_segments: tuple[int, ...],
device: torch.device,
exempt: bool = True,
dense_layers: tuple[int, ...] = (),
tile_size: int = _TILE_ELEMS,
**kwargs: dict[str, Any],
) -> MiniMaxH3VSAMetadata:
tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
if tile_shape is None:
raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
raw_latent_shape[2] // patch_size[2])
prefix_segments = tuple(int(s) for s in prefix_segments if s > 0)
total_seq_length = sum(prefix_segments) + math.prod(dit_seq_shape)
(_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles,
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
return MiniMaxH3VSAMetadata(
current_timestep=current_timestep,
@@ -186,13 +287,14 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
exempt=exempt,
variable_block_sizes=variable_block_sizes,
untile_combined_index=untile_combined_index,
tile_elems=int(tile_size),
dense_layers=tuple(int(layer) for layer in dense_layers),
tile_buf_holder=self._tile_buf_holder,
)
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems: int = _TILE_ELEMS) -> torch.Tensor:
"""fp32 mean over each tile_elems-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
written), so a plain sum with fp32 accumulation needs no validity mask
@@ -200,8 +302,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Te
the masked mean exactly.
"""
batch, seq_len, heads, dim = x.shape
n_tiles = seq_len // _TILE_ELEMS
pooled = x.view(batch, n_tiles, _TILE_ELEMS, heads, dim).sum(dim=2, dtype=torch.float32)
n_tiles = seq_len // tile_elems
pooled = x.view(batch, n_tiles, tile_elems, heads, dim).sum(dim=2, dtype=torch.float32)
pooled = pooled / variable_block_sizes.view(1, -1, 1, 1)
return pooled.permute(0, 2, 1, 3)
@@ -232,6 +334,24 @@ def _build_block_mask(
return mask
def _sm100a_unavailable_reason(sm100a_mod: Any, query_bhsd: torch.Tensor, variable_block_sizes: torch.Tensor,
grad_mode: bool) -> str | None:
"""Why the opt-in sm_100a forward route cannot run here, or None if it can.
Pure decision logic, split out so the routing is unit-testable without a
GPU or the compiled extension (tests substitute ``sm100a_mod``). Order
matters only for the message: the cheapest, most actionable reason first.
"""
if sm100a_mod is None:
return "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
if grad_mode:
return "inputs require grad and the sm_100a kernel is forward-only; grad paths keep Triton"
if not sm100a_mod.is_supported(query_bhsd, variable_block_sizes):
return ("block_sparse_attn_sm100a.is_supported returned False (needs an sm_100 device, a built "
"extension, bf16, head_dim 128, an even tile count, and integer tile sizes)")
return None
class MiniMaxH3VSAImpl(AttentionImpl):
def __init__(
@@ -259,7 +379,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
"routed to the VSA-H3 backend; exclude it from the supported backends.")
n_tiles = attn_metadata.variable_block_sizes.numel()
target_shape = (x.shape[0], n_tiles * _TILE_ELEMS, x.shape[-2], x.shape[-1])
target_shape = (x.shape[0], n_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])
# single scatter: untile_combined_index maps original row i to its
# padded slot, so this is exactly the inverse of postprocess_output
@@ -281,7 +401,11 @@ class MiniMaxH3VSAImpl(AttentionImpl):
gate_compress: torch.Tensor | None,
attn_metadata: MiniMaxH3VSAMetadata,
) -> torch.Tensor:
if block_sparse_attn_256_bshd is None:
tile_elems = attn_metadata.tile_elems
if tile_elems == 64:
if block_sparse_attn_64_bhsd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn is not installed")
elif block_sparse_attn_256_bshd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
@@ -291,8 +415,8 @@ class MiniMaxH3VSAImpl(AttentionImpl):
scores = None
if layer_sparsity > 0.0 or gate_compress is not None or probe_dir is not None:
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (query.shape[-1]**0.5)
if probe_dir is not None:
record_probe(probe_dir, self.layer_idx, query, key, scores, attn_metadata)
@@ -309,18 +433,75 @@ class MiniMaxH3VSAImpl(AttentionImpl):
attn_metadata.exempt,
)
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
if tile_elems == 64:
# Native 64-token path: the block map is already at the kernels'
# granularity. Both 64-token entries take BHSD ([B, H, S_pad, D]);
# mirror block_sparse_attn_256_bshd's Triton branch and transpose
# around the call.
q_bhsd = query.transpose(1, 2).contiguous()
k_bhsd = key.transpose(1, 2).contiguous()
v_bhsd = value.transpose(1, 2).contiguous()
# Opt-in sm_100a CUDA forward (upstream PR #1719 + per-q-tile
# q2k_num fix). Forward-only: grad-tracking calls stay on Triton
# so autograd keeps the Triton fwd+bwd pairing untouched. The
# kernel does return an LSE in Triton's M format, so a future
# fwd/bwd pairing is possible, but it is not built here.
use_sm100a = False
if os.environ.get(VSA_SM100A_ENV, "0") == "1":
grad_mode = torch.is_grad_enabled() and (query.requires_grad or key.requires_grad
or value.requires_grad)
reason = _sm100a_unavailable_reason(_sm100a, q_bhsd, attn_metadata.variable_block_sizes, grad_mode)
if reason is None and map_to_index is None:
reason = "fastvideo_kernel.triton_kernels.index (map_to_index) is not importable"
if reason is None:
use_sm100a = True
elif not torch.compiler.is_compiling():
logger.warning_once(f"{VSA_SM100A_ENV}=1 but falling back to the Triton-64 kernels: {reason}")
if use_sm100a:
# The sm_100a entry is index-native; compact the bool map the
# same way the Triton bool entry does internally. Per-row
# counts are NON-uniform here (prefix query tiles are dense,
# video tiles run prefix+top-k) -- legal for the fixed kernel,
# silently wrong on the pre-fix upstream one.
q2k_idx, q2k_num = map_to_index(mask)
out_bhsd, _ = _sm100a.block_sparse_attn_sm100a(
q_bhsd,
k_bhsd,
v_bhsd,
q2k_idx,
q2k_num,
attn_metadata.variable_block_sizes.to(torch.int32),
need_lse=False,
)
else:
out_bhsd, _ = block_sparse_attn_64_bhsd(
q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
)
out = out_bhsd.transpose(1, 2).contiguous()
else:
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
if gate_compress is not None:
# Wan-style compression branch: dense attention over pooled tiles,
# broadcast to each tile's rows, scaled by the learned gate
# (zero-initialized for H3 => branch contributes nothing until
# finetuned; the model layer skips it entirely for all-zero gates).
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
batch, _, heads, dim = out.shape
batch, seq_len, heads, dim = out.shape
n_tiles = attn_metadata.variable_block_sizes.numel()
out.view(batch, n_tiles, _TILE_ELEMS, heads,
dim).addcmul_(out_c.unsqueeze(2), gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim))
# Out-of-place: on the CuTe backend ``out`` is the tensor FA4's
# autograd node saved for its backward, so an in-place add here
# bumps its version counter and backward dies with "one of the
# variables needed for gradient computation has been modified".
out_tiled = out.view(batch, n_tiles, tile_elems, heads, dim)
gate_tiled = gate_compress.view(batch, n_tiles, tile_elems, heads, dim)
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
return out
@@ -60,7 +60,7 @@ def record_probe(
gen = torch.Generator(device="cpu").manual_seed(step * 1000 + layer)
# sample among video rows in the PADDED/tiled domain that are non-pad
from fastvideo.attention.backends.video_sparse_attn_h3 import token_tile_and_valid
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes)
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
video_rows = torch.nonzero((token_tile >= P) & token_valid, as_tuple=False).flatten()
idx = video_rows[torch.randint(0, video_rows.numel(), (_TRUE_ROWS, ), generator=gen).to(query.device)]
+1
View File
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
+375
View File
@@ -0,0 +1,375 @@
# SPDX-License-Identifier: Apache-2.0
"""Benchmark generate-fewer-frames plus MLX RIFE interpolation.
This script keeps the video diffusion path delegated to
``fastvideo.benchmarks.mlx_fastwan_bench._generate_cell``. It only orchestrates
two frame-count cells and the postprocess interpolation step.
"""
from __future__ import annotations
import argparse
import json
import shutil
import time
from pathlib import Path
from types import SimpleNamespace
import imageio.v2 as imageio
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import encode_prompt, make_rotary_embeddings
from fastvideo.benchmarks.mlx_fastwan_bench import _generate_cell, _ms_ssim
from fastvideo.mlx_runtime.rife_interp import RIFEBackendError, interpolate, load_model
from fastvideo.mlx_runtime.memory import add_memory_limit_args, apply_memory_limits
FOX_PROMPT = "A fox runs through a misty pine forest, leaves kicking up behind it."
DEFAULT_MODEL_ROOT = Path("/Users/aryank/models/qad_int8_v2")
def _parse_timesteps(raw: str) -> list[int]:
timesteps = [int(part.strip()) for part in raw.split(",") if part.strip()]
if not timesteps:
raise SystemExit("No DMD timesteps parsed from --dmd-denoising-steps")
return timesteps
def _read_video(path: Path) -> list[np.ndarray]:
if not path.is_file():
raise FileNotFoundError(f"Video does not exist: {path}")
frames = [np.asarray(frame[:, :, :3], dtype=np.uint8) for frame in imageio.mimread(path)]
if not frames:
raise RuntimeError(f"No frames decoded from {path}")
return frames
def _write_video(path: Path, frames: list[np.ndarray], fps: int) -> None:
if not frames:
raise ValueError("Cannot write an empty frame list")
path.parent.mkdir(parents=True, exist_ok=True)
with imageio.get_writer(str(path), fps=fps, macro_block_size=1, codec="libx264", quality=8) as writer:
for frame in frames:
writer.append_data(np.asarray(frame, dtype=np.uint8))
def _copy_video(src: Path, dst: Path) -> Path:
dst.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(src, dst)
return dst
def _make_generation_inputs(args, num_frames: int):
import mlx.core as mx
import torch
config_path = args.model_root / "transformer" / "config.json"
checkpoint_path = args.model_root / "transformer" / "diffusion_pytorch_model.safetensors"
if not config_path.is_file():
raise SystemExit(f"Missing DiT config: {config_path}")
if not checkpoint_path.is_file():
raise SystemExit(f"Missing DiT checkpoint: {checkpoint_path}")
config = json.loads(config_path.read_text())
latent_frames = (num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
freqs_cis = make_rotary_embeddings(
config,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
)
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()
timesteps = _parse_timesteps(args.dmd_denoising_steps)
renoise_by_step = [
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
for _ in range(max(0,
len(timesteps) - 1))
]
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,
)
return {
"checkpoint_path": checkpoint_path,
"config_path": config_path,
"encoder_hidden_states": mx.array(prompt_embeds.numpy()),
"freqs_cis": freqs_cis,
"timesteps": timesteps,
"latents_seed": latents_seed,
"renoise_by_step": renoise_by_step,
"latent_frames": latent_frames,
}
def _generate(args, num_frames: int):
cell_args = SimpleNamespace(
model_root=args.model_root,
height=args.height,
width=args.width,
num_frames=num_frames,
fps=args.fps,
flow_shift=args.flow_shift,
torch_device=args.torch_device,
torch_dtype=args.torch_dtype,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
mlx_checkpoint_cache=args.mlx_checkpoint_cache,
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
mlx_disable_cache=args.mlx_disable_cache,
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
benchmark_preset="metalfx-rife",
current_prompt_id=f"fox-{num_frames}f",
current_prompt=args.prompt,
output_dir=args.output_dir,
)
inputs = _make_generation_inputs(args, num_frames)
print(f"=== generate: {num_frames} frames ({inputs['latent_frames']} latent frames) ===", flush=True)
return _generate_cell(
args=cell_args,
mode=args.mode,
decoder=args.decoder,
checkpoint_path=inputs["checkpoint_path"],
config_path=inputs["config_path"],
encoder_hidden_states=inputs["encoder_hidden_states"],
freqs_cis=inputs["freqs_cis"],
timesteps=inputs["timesteps"],
latents_seed=inputs["latents_seed"],
renoise_by_step=inputs["renoise_by_step"],
)
def _relative_to_output(path: Path, output_dir: Path) -> str:
try:
return str(path.relative_to(output_dir))
except ValueError:
return str(path)
def _write_report(args, result: dict) -> Path:
report_path = args.output_dir / "metalfx_rife_report.md"
rows = result["speed_rows"]
table_lines = [
"| path | denoise_s | decode_s | gen_total_s | rife_s | net_s | speedup_vs_81 |",
"| --- | ---: | ---: | ---: | ---: | ---: | ---: |",
]
for row in rows:
table_lines.append(
"| {path} | {denoise_s:.3f} | {decode_s:.3f} | {gen_total_s:.3f} | {rife_s:.3f} | {net_s:.3f} | {speedup_vs_81:.3f}x |"
.format(**row))
text = f"""# MetalFX/RIFE Generate-Fewer-Frames Benchmark Run
Prompt: {args.prompt}
Resolution: {args.height}x{args.width}, fps={args.fps}, mode={args.mode}, decoder={args.decoder}
| metric | value |
| --- | ---: |
| reconstruction_ms_ssim | {result['reconstruction_ms_ssim']:.6f} |
| reference_frames | {result['reference_frames']} |
| reduced_frames | {result['reduced_frames']} |
| interpolated_frames | {result['interpolated_frames']} |
{chr(10).join(table_lines)}
Videos:
- reference: `{result['videos']['reference']}`
- reference drop-41 RIFE reconstruction: `{result['videos']['drop41_rife81']}`
- generated 41 RIFE to 81: `{result['videos']['generated41_rife81']}`
- generated 41 direct: `{result['videos']['generated41']}`
Raw metrics are in `{_relative_to_output(args.output_dir / 'metrics.json', args.output_dir)}`.
"""
report_path.write_text(text)
return report_path
def main() -> None:
parser = argparse.ArgumentParser(
description="Evaluate 41-frame generation plus MLX RIFE interpolation vs 81-frame generation.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--prompt", default=FOX_PROMPT)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=832)
parser.add_argument("--reference-frames", type=int, default=81)
parser.add_argument("--reduced-frames", type=int, default=41)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--mode", default="int8", choices=("fp16", "bf16", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"))
parser.add_argument("--decoder", default="taehv", choices=("taehv", "wan-vae"))
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=512)
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--output-dir", type=Path, default=Path("bench/metalfx_rife"))
parser.add_argument("--mlx-checkpoint-cache", type=Path, default=None)
parser.add_argument("--compile", action="store_true", help="Enable FASTVIDEO_MLX_COMPILE=1 for DiT denoise.")
parser.add_argument("--rife-scale", type=float, default=1.0)
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")
add_memory_limit_args(parser)
args = parser.parse_args()
if args.reference_frames != (args.reduced_frames - 1) * 2 + 1:
raise SystemExit(
"--reference-frames must equal (--reduced-frames - 1) * 2 + 1 for the default every-other-frame test")
if args.compile:
import os
os.environ["FASTVIDEO_MLX_COMPILE"] = "1"
args.model_root = args.model_root.expanduser().resolve()
args.output_dir = args.output_dir.expanduser().resolve()
args.output_dir.mkdir(parents=True, exist_ok=True)
if args.mlx_checkpoint_cache is None:
args.mlx_checkpoint_cache = args.output_dir / "mlx_checkpoint_cache"
else:
args.mlx_checkpoint_cache = args.mlx_checkpoint_cache.expanduser().resolve()
import mlx.core as mx
import torch
runtime_limits = apply_memory_limits(
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
mlx_disable_cache=args.mlx_disable_cache,
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
mx_module=mx,
).as_metrics()
mx.random.seed(args.seed)
torch.manual_seed(args.seed)
reference_cell = _generate(args, args.reference_frames)
reduced_cell = _generate(args, args.reduced_frames)
reference_video = _copy_video(reference_cell.video_path, args.output_dir / "fox_reference_81.mp4")
generated41_video = _copy_video(reduced_cell.video_path, args.output_dir / "fox_generated_41.mp4")
reference_frames = _read_video(reference_video)
dropped_reference_frames = reference_frames[::2]
if len(dropped_reference_frames) != args.reduced_frames:
raise RuntimeError(f"Expected {args.reduced_frames} dropped frames, got {len(dropped_reference_frames)}")
try:
rife_model = load_model("4.25")
except RIFEBackendError:
raise
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
raise RIFEBackendError(f"Unexpected RIFE load failure: {exc}") from exc
print("=== RIFE reconstruction: drop every other reference frame, 41 -> 81 ===", flush=True)
start = time.perf_counter()
reconstructed_frames = interpolate(dropped_reference_frames, factor=2, model=rife_model, scale=args.rife_scale)
recon_rife_s = time.perf_counter() - start
reconstructed_video = args.output_dir / "fox_reference_drop41_rife81.mp4"
_write_video(reconstructed_video, reconstructed_frames, args.fps)
reconstruction_ms_ssim = _ms_ssim(reference_video, reconstructed_video, required=True)
if reconstruction_ms_ssim is None:
raise RuntimeError("MS-SSIM returned None for reconstruction comparison")
print("=== RIFE actual reduced path: generated 41 -> 81 ===", flush=True)
generated41_frames = _read_video(generated41_video)
start = time.perf_counter()
generated41_rife_frames = interpolate(generated41_frames, factor=2, model=rife_model, scale=args.rife_scale)
generated41_rife_s = time.perf_counter() - start
generated41_rife_video = args.output_dir / "fox_generated41_rife81.mp4"
_write_video(generated41_rife_video, generated41_rife_frames, args.fps)
direct_denoise_s = float(reference_cell.metrics["denoise_s"])
reduced_denoise_s = float(reduced_cell.metrics["denoise_s"])
direct_gen_total_s = float(reference_cell.metrics["denoise_s"]) + float(reference_cell.metrics["decode_s"])
reduced_gen_total_s = float(reduced_cell.metrics["denoise_s"]) + float(reduced_cell.metrics["decode_s"])
net_denoise_rife_s = reduced_denoise_s + generated41_rife_s
net_gen_rife_s = reduced_gen_total_s + generated41_rife_s
result = {
"prompt":
args.prompt,
"model_root":
str(args.model_root),
"rife_impl":
"rife-mlx vendored at fastvideo/third_party/rife_mlx, weights mlx-community/RIFE-4.25",
"runtime_limits":
runtime_limits,
"reference_frames":
len(reference_frames),
"reduced_frames":
len(generated41_frames),
"interpolated_frames":
len(generated41_rife_frames),
"reconstruction_ms_ssim":
reconstruction_ms_ssim,
"reconstruction_rife_s":
recon_rife_s,
"generated41_rife_s":
generated41_rife_s,
"reference_metrics":
reference_cell.metrics,
"reduced_metrics":
reduced_cell.metrics,
"speed_rows": [
{
"path": "generate_81",
"denoise_s": direct_denoise_s,
"decode_s": float(reference_cell.metrics["decode_s"]),
"gen_total_s": direct_gen_total_s,
"rife_s": 0.0,
"net_s": direct_gen_total_s,
"speedup_vs_81": 1.0,
},
{
"path": "generate_41_plus_rife81_denoise_only",
"denoise_s": reduced_denoise_s,
"decode_s": 0.0,
"gen_total_s": reduced_denoise_s,
"rife_s": generated41_rife_s,
"net_s": net_denoise_rife_s,
"speedup_vs_81": direct_denoise_s / net_denoise_rife_s,
},
{
"path": "generate_41_plus_rife81_decode_included",
"denoise_s": reduced_denoise_s,
"decode_s": float(reduced_cell.metrics["decode_s"]),
"gen_total_s": reduced_gen_total_s,
"rife_s": generated41_rife_s,
"net_s": net_gen_rife_s,
"speedup_vs_81": direct_gen_total_s / net_gen_rife_s,
},
],
"videos": {
"reference": _relative_to_output(reference_video, args.output_dir),
"drop41_rife81": _relative_to_output(reconstructed_video, args.output_dir),
"generated41": _relative_to_output(generated41_video, args.output_dir),
"generated41_rife81": _relative_to_output(generated41_rife_video, args.output_dir),
},
}
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(result, indent=2))
report_path = _write_report(args, result)
print(json.dumps(result["speed_rows"], indent=2))
print(f"reconstruction_ms_ssim={reconstruction_ms_ssim:.6f}")
print(f"wrote {metrics_path}")
print(f"wrote {report_path}")
if __name__ == "__main__":
main()
+826
View File
@@ -0,0 +1,826 @@
# SPDX-License-Identifier: Apache-2.0
"""Prove-out benchmark for the MLX FastWan runtime (Apple Silicon).
Sweeps ``{dtype/quant} x {decoder}``, generates a clip per cell, and records the
latency breakdown, peak unified memory, and MS-SSIM (optionally LPIPS) against a
reference video. It emits a JSON blob and a markdown table -- the artifact that
turns "int8 + TAEHV looks good" into defensible numbers, and (via
``--assert-min-ssim``) a regression gate for the ``mx.compile`` work.
Design notes:
- Generation reuses the hybrid POC helpers in
``examples/inference/basic/mlx_wan_prompt_to_video.py`` (torch-MPS UMT5 encode
and Wan-VAE/TAEHV decode) plus the on-device MLX DMD sampler
(``fastvideo/mlx_runtime/sampling.py``); the denoise loop never leaves the
device.
- Quality reuses the tested MS-SSIM primitive
``fastvideo/tests/utils.py::compute_video_ssim_torchvision``.
- Reference: by default each cell is scored against the highest-fidelity cell
in the sweep (``fp16`` + ``wan-vae``), which needs no CUDA box and answers
"how much does int8/int4/TAEHV degrade vs the best local config". Pass
``--reference PATH`` to score against an external clip instead (e.g. the
torch-MPS or CUDA FastVideo output of the same model) for a "vs. the original
model" column.
Run on an Apple Silicon Mac (needs ``mlx`` + a torch build with MPS):
python fastvideo/benchmarks/mlx_fastwan_bench.py \
--modes fp16,bf16,int8,int4 --decoders taehv,wan-vae
"""
from __future__ import annotations
import argparse
import html
import json
import os
import time
from dataclasses import dataclass, field
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,
encode_prompt,
make_rotary_embeddings,
)
from fastvideo.mlx_runtime.memory import add_memory_limit_args, apply_memory_limits, cleanup_mlx
# The highest-fidelity cell; used as the default SSIM reference when no external
# reference video is supplied.
REFERENCE_MODE = "fp16"
REFERENCE_DECODER = "wan-vae"
ALLOWED_MODES = ("fp16", "bf16", "int8", "int4", "mxfp8", "mxfp4", "nvfp4")
ALLOWED_DECODERS = ("taehv", "wan-vae")
@dataclass(frozen=True)
class PromptCase:
id: str
prompt: str
@dataclass(frozen=True)
class BenchmarkPreset:
height: int
width: int
num_frames: int
modes: str
decoders: str
mlx_memory_limit_gib: float | None = None
mlx_disable_cache: bool = False
torch_mps_high_watermark_ratio: float | None = None
torch_mps_low_watermark_ratio: float | None = None
PROMPT_SETS = {
"motion7": (
PromptCase("beach-sunset", "A slow cinematic sunset over ocean waves at a quiet beach."),
PromptCase("fox-forest", "A fox runs through a misty pine forest, leaves kicking up behind it."),
PromptCase("raccoon-sunflowers", "A raccoon walks through a sunflower field as petals move in the wind."),
PromptCase("surfing-cat", "A cat wearing sunglasses surfs across a bright blue ocean wave."),
PromptCase("burning-clock", "A vintage table clock burns on a wooden desk, flames flickering realistically."),
PromptCase("forest-walk", "Video game style footage of a man walking through a dense forest path."),
PromptCase("sea-dock-yachts", "Several yachts are parked at a sea dock while water ripples around them."),
),
}
BENCHMARK_PRESETS = {
"default":
BenchmarkPreset(
height=480,
width=832,
num_frames=81,
modes="fp16,bf16,int8,int4",
decoders="taehv,wan-vae",
),
"mac-16gb":
BenchmarkPreset(
height=448,
width=832,
num_frames=61,
modes="int8",
decoders="taehv",
mlx_memory_limit_gib=16.0,
mlx_disable_cache=True,
torch_mps_high_watermark_ratio=0.57,
torch_mps_low_watermark_ratio=0.0,
),
"mac-32gb":
BenchmarkPreset(
height=480,
width=832,
num_frames=81,
modes="int8,fp16",
decoders="taehv",
),
"mac-64gb":
BenchmarkPreset(
height=480,
width=832,
num_frames=81,
modes="int8,fp16",
decoders="taehv,wan-vae",
),
}
@dataclass
class Cell:
prompt_id: str
prompt: str
mode: str
decoder: str
video_path: Path
latents: np.ndarray
metrics: dict[str, float | int | str | bool | None] = field(default_factory=dict)
def _mode_to_dtype_quant(mode: str) -> tuple[str, str | None]:
"""Map a sweep mode to (MLX compute dtype, quantization spec).
Quantized modes keep fp16 activations and quantize only the DiT linear
weights (matching ``mlx_dit_from_diffusers_safetensors``).
"""
if mode == "bf16":
return "bf16", None
if mode == "fp16":
return "fp16", None
# int8/int4/mxfp*/nvfp4 -> fp16 activations + quantized weights.
return "fp16", mode
def _mx_dtype(mx, base: str):
return {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[base]
def _parse_list(raw: str, allowed: tuple[str, ...], label: str) -> list[str]:
items = [x.strip() for x in raw.split(",") if x.strip()]
unknown = sorted(set(items) - set(allowed))
if unknown:
raise ValueError(f"Unsupported {label}: {unknown} (allowed: {list(allowed)})")
return items
def _safe_slug(value: str, *, fallback: str) -> str:
slug = "".join(ch.lower() if ch.isalnum() else "-" for ch in value.strip())
slug = "-".join(part for part in slug.split("-") if part)
return slug[:64] or fallback
def _load_prompt_cases(prompt: str, prompt_file: Path | None, prompt_set: str = "single") -> list[PromptCase]:
"""Load one prompt, a built-in prompt set, or a text/jsonl prompt file."""
if prompt_file is None:
if prompt_set == "single":
return [PromptCase(id="prompt-001", prompt=prompt)]
if prompt_set not in PROMPT_SETS:
raise ValueError(f"Unsupported prompt set: {prompt_set} (allowed: {sorted(PROMPT_SETS) + ['single']})")
return list(PROMPT_SETS[prompt_set])
cases: list[PromptCase] = []
for line_index, raw_line in enumerate(prompt_file.read_text().splitlines(), start=1):
line = raw_line.strip()
if not line or line.startswith("#"):
continue
prompt_id = f"prompt-{len(cases) + 1:03d}"
prompt_text = line
if prompt_file.suffix.lower() == ".jsonl":
item = json.loads(line)
prompt_text = str(item.get("prompt") or item.get("text") or item.get("caption") or "").strip()
if not prompt_text:
raise ValueError(f"{prompt_file}:{line_index} has no prompt/text/caption field")
prompt_id = str(item.get("id") or item.get("name") or prompt_id)
cases.append(PromptCase(id=_safe_slug(prompt_id, fallback=f"prompt-{len(cases) + 1:03d}"), prompt=prompt_text))
if not cases:
raise ValueError(f"No prompts found in {prompt_file}")
return cases
def denoise_dmd_on_device(
*,
mx,
dit,
latents,
encoder_hidden_states,
freqs_cis,
timesteps: list[int],
renoise_by_step: list[np.ndarray],
schedule,
dmd_step,
mx_dtype,
) -> tuple[np.ndarray, list[float]]:
"""Run the FastWan DMD loop entirely on the MLX device.
Mirrors the loop in ``mlx_wan_prompt_to_video.py`` (fp32 affine math, MLX RNG
re-noise) so the benchmark measures exactly the shipped path.
Returns the final latents plus per-step wall times. The first step carries
one-time costs (mx.compile tracing, kernel warm-up), so first-vs-steady
step timing is how the benchmark separates cold-start from steady-state
denoise throughput.
All host-side tensors (timesteps, re-noise draws) are uploaded before the
loop starts, so the per-step body performs no bulk host->device transfers
and step timings measure device work rather than staging copies.
"""
timesteps_mx = [mx.array([float(timestep)]).astype(mx.float32) for timestep in timesteps]
renoise_mx = [mx.array(renoise).astype(mx.float32) for renoise in renoise_by_step]
if timesteps_mx or renoise_mx:
mx.eval(*timesteps_mx, *renoise_mx)
step_times: list[float] = []
for step_index, timestep in enumerate(timesteps):
step_start = time.perf_counter()
noise_input_latent = latents
noise_pred = dit(latents.astype(mx_dtype), encoder_hidden_states, timesteps_mx[step_index], freqs_cis)
noise_input_f32 = noise_input_latent.astype(mx.float32)
pred_noise_f32 = noise_pred.astype(mx.float32)
if step_index < len(timesteps) - 1:
next_ts: float | None = float(timesteps[step_index + 1])
renoise = renoise_mx[step_index]
else:
next_ts, renoise = None, None
latents = dmd_step(
latents=noise_input_f32,
noise_input_latent=noise_input_f32,
pred_noise=pred_noise_f32,
schedule=schedule,
timestep=float(timestep),
next_timestep=next_ts,
noise=renoise,
).astype(mx_dtype)
mx.eval(latents)
step_times.append(time.perf_counter() - step_start)
return np.array(latents.astype(mx.float32)), step_times
def _peak_memory_bytes(mx) -> int:
try:
return int(mx.get_peak_memory())
except Exception: # noqa: BLE001 - best-effort telemetry only.
return 0
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
mse = float(np.mean(np.square(diff)))
signal = float(np.mean(np.square(baseline.astype(np.float32))))
return {
"latent_mse_vs_ref_mode": mse,
"latent_snr_db_vs_ref_mode": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
}
def _ms_ssim(reference_video: Path, candidate_video: Path, *, required: bool = False) -> float | None:
"""Mean MS-SSIM between two mp4s, via the repo's tested helper."""
if not reference_video.exists() or not candidate_video.exists():
return None
try:
from fastvideo.tests.utils import compute_video_ssim_torchvision
except ImportError as exc:
message = ("MS-SSIM is unavailable because `pytorch-msssim` is not installed. "
"Install FastVideo with the test extra, e.g. `uv pip install -e '.[mlx,test]'`, "
"or run without an SSIM assertion.")
if required:
raise RuntimeError(message) from exc
print(f"{message} Skipping MS-SSIM.")
return None
try:
ssim_values = compute_video_ssim_torchvision(str(reference_video), str(candidate_video), use_ms_ssim=True)
except ImportError as exc:
message = ("MS-SSIM is unavailable because `pytorch-msssim` is not installed. "
"Install FastVideo with the test extra, e.g. `uv pip install -e '.[mlx,test]'`, "
"or run without an SSIM assertion.")
if required:
raise RuntimeError(message) from exc
print(f"{message} Skipping MS-SSIM.")
return None
return float(ssim_values[0])
def _markdown_table(rows: list[dict]) -> str:
columns = [
("prompt_id", "prompt"),
("mode", "mode"),
("decoder", "decoder"),
("status", "status"),
("denoise_s", "denoise s"),
("decode_s", "decode s"),
("total_s", "total s"),
("peak_gib", "peak GiB"),
("ms_ssim_vs_ref", "MS-SSIM"),
("lpips_vs_ref", "LPIPS"),
]
header = "| " + " | ".join(label for _, label in columns) + " |"
sep = "| " + " | ".join("---" for _ in columns) + " |"
lines = [header, sep]
for row in rows:
cells = []
for key, _ in columns:
value = row.get(key)
if isinstance(value, float):
cells.append(f"{value:.3f}")
elif value is None:
cells.append("-")
else:
cells.append(str(value))
lines.append("| " + " | ".join(cells) + " |")
return "\n".join(lines)
def _format_metric(value) -> str:
if isinstance(value, float):
return f"{value:.3f}"
if value is None:
return "-"
return str(value)
def _html_grid(rows: list[dict]) -> str:
groups: dict[str, list[dict]] = {}
for row in rows:
groups.setdefault(str(row.get("prompt_id", "prompt")), []).append(row)
sections = []
for prompt_id, group_rows in groups.items():
prompt = next((str(row.get("prompt", "")) for row in group_rows if row.get("prompt")), "")
cards = []
for row in group_rows:
title = f"{row.get('mode', '-')} / {row.get('decoder', '-')}"
status = row.get("status", "-")
video_path = row.get("video_path")
if video_path:
media = f'<video src="{html.escape(str(video_path))}" muted loop controls playsinline></video>'
else:
media = f'<div class="missing">No video<br>{html.escape(str(row.get("error", "")))}</div>'
metrics = (f"status={status} · total={_format_metric(row.get('total_s'))}s · "
f"denoise={_format_metric(row.get('denoise_s'))}s · "
f"decode={_format_metric(row.get('decode_s'))}s · "
f"peak={_format_metric(row.get('peak_gib'))}GiB")
cards.append("<article>"
f"<h3>{html.escape(title)}</h3>"
f"{media}"
f"<p>{html.escape(metrics)}</p>"
"</article>")
sections.append("<section>"
f"<h2>{html.escape(prompt_id)}</h2>"
f"<p class=\"prompt\">{html.escape(prompt)}</p>"
f"<div class=\"grid\">{''.join(cards)}</div>"
"</section>")
return """<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>FastVideo MLX benchmark grid</title>
<style>
body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif; margin: 24px; background: #111; color: #eee; }
button { margin-right: 8px; padding: 8px 12px; border-radius: 8px; border: 1px solid #555; background: #222; color: #eee; }
section { margin-top: 28px; }
.prompt { color: #bbb; max-width: 900px; }
.grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(280px, 1fr)); gap: 16px; }
article { background: #1b1b1b; border: 1px solid #333; border-radius: 12px; padding: 12px; }
h1, h2, h3 { margin: 0 0 10px; }
video { width: 100%; border-radius: 8px; background: #000; }
article p { color: #bbb; font-size: 13px; line-height: 1.4; }
.missing { min-height: 160px; display: grid; place-items: center; text-align: center; color: #f5b5b5; background: #2a1515; border-radius: 8px; padding: 12px; }
</style>
</head>
<body>
<h1>FastVideo MLX benchmark grid</h1>
<p>Use the controls below to start/stop every clip together for side-by-side inspection.</p>
<button onclick="for (const v of document.querySelectorAll('video')) { v.currentTime = 0; v.play(); }">Restart + play all</button>
<button onclick="for (const v of document.querySelectorAll('video')) v.pause();">Pause all</button>
""" + "\n".join(sections) + """
</body>
</html>
"""
def _write_html_grid(rows: list[dict], output_dir: Path) -> Path:
html_path = output_dir / "index.html"
html_path.write_text(_html_grid(rows))
return html_path
def _generate_cell(
*,
args,
mode: str,
decoder: str,
checkpoint_path: Path,
config_path: Path,
encoder_hidden_states,
freqs_cis,
timesteps: list[int],
latents_seed: np.ndarray,
renoise_by_step: list[np.ndarray],
) -> Cell:
import mlx.core as mx
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
base_dtype, quantization = _mode_to_dtype_quant(mode)
mx_dtype = _mx_dtype(mx, base_dtype)
mx.clear_cache()
mx.reset_peak_memory()
load_start = time.perf_counter()
load_source = "diffusers"
if args.mlx_checkpoint_cache is not None:
from fastvideo.mlx_runtime.checkpoint import (
load_mlx_dit_checkpoint,
save_mlx_dit_checkpoint,
)
mode_ckpt_dir = args.mlx_checkpoint_cache / mode
if (mode_ckpt_dir / "mlx_dit.json").exists():
dit = load_mlx_dit_checkpoint(mode_ckpt_dir)
load_source = "mlx_checkpoint"
else:
dit = mlx_dit_from_diffusers_safetensors(
checkpoint_path,
config_path,
dtype=base_dtype,
quantization=quantization,
)
save_mlx_dit_checkpoint(dit, mode_ckpt_dir)
load_source = "diffusers_then_saved"
else:
dit = mlx_dit_from_diffusers_safetensors(
checkpoint_path,
config_path,
dtype=base_dtype,
quantization=quantization,
)
load_s = time.perf_counter() - load_start
load_peak = _peak_memory_bytes(mx)
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
latents = mx.array(latents_seed).astype(mx_dtype)
mx.reset_peak_memory()
denoise_start = time.perf_counter()
latents_np, step_times = denoise_dmd_on_device(
mx=mx,
dit=dit,
latents=latents,
encoder_hidden_states=encoder_hidden_states.astype(mx_dtype),
freqs_cis=freqs_cis,
timesteps=timesteps,
renoise_by_step=renoise_by_step,
schedule=schedule,
dmd_step=dmd_step,
mx_dtype=mx_dtype,
)
denoise_s = time.perf_counter() - denoise_start
denoise_peak = _peak_memory_bytes(mx)
del dit, latents
cleanup_mlx(mx)
video_path = (args.output_dir / f"{args.current_prompt_id}" /
f"video_{mode}_{decoder}_{args.height}x{args.width}x{args.num_frames}.mp4")
decode_start = time.perf_counter()
decode_latents_to_video(
model_root=args.model_root,
latents_np=latents_np,
output_path=video_path,
fps=args.fps,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
backend=decoder,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
)
decode_s = time.perf_counter() - decode_start
metrics: dict[str, float | int | str | bool | None] = {
"prompt_id": args.current_prompt_id,
"prompt": args.current_prompt,
"benchmark_preset": args.benchmark_preset,
"mode": mode,
"decoder": decoder,
"status": "ok",
"video_path": str(video_path.relative_to(args.output_dir)),
"load_s": load_s,
"load_source": load_source,
"denoise_s": denoise_s,
# The first step carries one-time costs (mx.compile tracing, kernel
# warm-up); steady-state throughput is the median of the rest.
"denoise_first_step_s": step_times[0] if step_times else None,
"denoise_steady_step_s": (float(np.median(step_times[1:])) if len(step_times) > 1 else None),
"decode_s": decode_s,
"total_s": load_s + denoise_s + decode_s,
"load_peak_gib": load_peak / (1024**3),
"peak_gib": max(load_peak, denoise_peak) / (1024**3),
"quantization": quantization or "none",
"compute_dtype": base_dtype,
"compile": os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1",
"fast_norm": os.environ.get("FASTVIDEO_MLX_FAST_NORM", "0") == "1",
"mlx_memory_limit_gib": args.mlx_memory_limit_gib,
"mlx_cache_limit_gib": args.mlx_cache_limit_gib,
"mlx_disable_cache": args.mlx_disable_cache,
"mlx_wired_limit_gib": args.mlx_wired_limit_gib,
"torch_mps_high_watermark_ratio": args.torch_mps_high_watermark_ratio,
"torch_mps_low_watermark_ratio": args.torch_mps_low_watermark_ratio,
}
return Cell(
prompt_id=args.current_prompt_id,
prompt=args.current_prompt,
mode=mode,
decoder=decoder,
video_path=video_path,
latents=latents_np,
metrics=metrics,
)
def main() -> None:
preset_parser = argparse.ArgumentParser(add_help=False)
preset_parser.add_argument("--benchmark-preset", choices=tuple(BENCHMARK_PRESETS), default="default")
preset_args, _ = preset_parser.parse_known_args()
preset = BENCHMARK_PRESETS[preset_args.benchmark_preset]
parser = argparse.ArgumentParser(description="MLX FastWan prove-out benchmark (latency + quality).")
parser.add_argument("--benchmark-preset",
choices=tuple(BENCHMARK_PRESETS),
default=preset_args.benchmark_preset,
help="Memory-tier benchmark defaults. Explicit CLI flags override preset values.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--prompt", default="A paper boat sails through a shallow stream in a mossy forest.")
parser.add_argument(
"--prompt-file",
type=Path,
default=None,
help=
"Optional prompt set. Plain text uses one prompt per non-empty line; .jsonl accepts prompt/text/caption plus optional id/name.",
)
parser.add_argument(
"--prompt-set",
choices=("single", *PROMPT_SETS.keys()),
default="single",
help="Built-in standard prompt set. Ignored when --prompt-file is supplied.",
)
parser.add_argument("--height", type=int, default=preset.height)
parser.add_argument("--width", type=int, default=preset.width)
parser.add_argument("--num-frames", type=int, default=preset.num_frames)
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=512)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--modes", default=preset.modes)
parser.add_argument("--decoders", default=preset.decoders)
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_fastwan_bench"))
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument(
"--reference",
type=Path,
default=None,
help="External reference mp4 to score every cell against. Defaults to the fp16+wan-vae cell.",
)
parser.add_argument("--assert-min-ssim",
type=float,
default=None,
help="Fail if any cell's MS-SSIM vs the reference falls below this value.")
parser.add_argument("--compile",
action="store_true",
help="Enable mx.compile on the DiT forward (sets FASTVIDEO_MLX_COMPILE=1).")
parser.add_argument("--lpips", action="store_true", help="Also compute LPIPS (needs the `lpips` package).")
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")
parser.add_argument(
"--mlx-checkpoint-cache",
type=Path,
default=None,
help="Directory of per-mode pre-quantized MLX checkpoints. The first cell of a mode "
"converts from Diffusers weights and saves here (load_source=diffusers_then_saved); "
"later cells and later runs reload without requantizing (load_source=mlx_checkpoint), "
"which is also how the checkpoint load-time win is measured.",
)
add_memory_limit_args(
parser,
mlx_memory_limit_gib=preset.mlx_memory_limit_gib,
mlx_disable_cache=preset.mlx_disable_cache,
torch_mps_high_watermark_ratio=preset.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=preset.torch_mps_low_watermark_ratio,
)
args = parser.parse_args()
if args.compile:
os.environ["FASTVIDEO_MLX_COMPILE"] = "1"
import mlx.core as mx
runtime_limits = apply_memory_limits(
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
mlx_disable_cache=args.mlx_disable_cache,
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
mx_module=mx,
).as_metrics()
import torch
mx.random.seed(args.seed)
torch.manual_seed(args.seed)
args.output_dir.mkdir(parents=True, exist_ok=True)
modes = _parse_list(args.modes, ALLOWED_MODES, "modes")
decoders = _parse_list(args.decoders, ALLOWED_DECODERS, "decoders")
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
freqs_cis = make_rotary_embeddings(
config,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
)
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()
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
# Keep DMD stochasticity identical across benchmark cells. Without this,
# FP16/INT8/decoder comparisons can accidentally measure different re-noise
# samples instead of only quantization or decoder differences.
renoise_by_step = [
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
for _ in range(max(0,
len(timesteps) - 1))
]
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
prompt_cases = _load_prompt_cases(args.prompt, args.prompt_file, args.prompt_set)
cells: list[Cell] = []
unsupported_rows: list[dict] = []
for prompt_case in prompt_cases:
args.current_prompt_id = prompt_case.id
args.current_prompt = prompt_case.prompt
prompt_embeds = encode_prompt(
model_root=args.model_root,
prompt=prompt_case.prompt,
max_sequence_length=args.max_sequence_length,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
)
encoder_hidden_states = mx.array(prompt_embeds.numpy())
for mode in modes:
for decoder in decoders:
print(f"=== cell: prompt={prompt_case.id} mode={mode} decoder={decoder} ===")
try:
cells.append(
_generate_cell(
args=args,
mode=mode,
decoder=decoder,
checkpoint_path=checkpoint_path,
config_path=config_path,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis,
timesteps=timesteps,
latents_seed=latents_seed,
renoise_by_step=renoise_by_step,
))
except UnsupportedMLXQuantizationError as exc:
# Record the cell as unsupported and keep sweeping: a partial
# report on this MLX build beats crashing the whole run.
print(f"skipping cell (unsupported by this MLX build): {exc}")
unsupported_rows.append({
"prompt_id": prompt_case.id,
"prompt": prompt_case.prompt,
"mode": mode,
"decoder": decoder,
"status": "unsupported_by_mlx",
"error": str(exc),
})
if not cells:
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(unsupported_rows, indent=2))
raise SystemExit(f"No benchmark cell could run: every requested mode is unsupported by this MLX build. "
f"Wrote {metrics_path}.")
# Resolve one internal reference per prompt. A single external reference, if
# supplied, is used for every prompt and only video metrics are computed.
reference_by_prompt: dict[str, tuple[Path, np.ndarray | None]] = {}
if args.reference is not None:
for prompt_case in prompt_cases:
reference_by_prompt[prompt_case.id] = (args.reference, None)
else:
for prompt_case in prompt_cases:
prompt_cells = [c for c in cells if c.prompt_id == prompt_case.id]
if not prompt_cells:
continue
ref_cell = next(
(c for c in prompt_cells if c.mode == REFERENCE_MODE and c.decoder == REFERENCE_DECODER),
prompt_cells[0],
)
reference_by_prompt[prompt_case.id] = (ref_cell.video_path, ref_cell.latents)
print(f"Using internal reference cell for {prompt_case.id}: "
f"mode={ref_cell.mode} decoder={ref_cell.decoder}")
lpips_fn = _load_lpips() if args.lpips else None
rows: list[dict] = []
failures: list[str] = []
for cell in cells:
reference_video, reference_latents = reference_by_prompt[cell.prompt_id]
ms_ssim = _ms_ssim(Path(reference_video), cell.video_path, required=args.assert_min_ssim is not None)
cell.metrics["ms_ssim_vs_ref"] = ms_ssim
cell.metrics.update(runtime_limits)
if reference_latents is not None:
cell.metrics.update(_latent_delta_metrics(cell.latents, reference_latents))
cell.metrics["lpips_vs_ref"] = (_lpips_between(lpips_fn, Path(reference_video), cell.video_path)
if lpips_fn else None)
if args.assert_min_ssim is not None and ms_ssim is not None and ms_ssim < args.assert_min_ssim:
failures.append(
f"{cell.prompt_id}/{cell.mode}/{cell.decoder}: MS-SSIM {ms_ssim:.4f} < {args.assert_min_ssim}")
rows.append(dict(cell.metrics))
print(json.dumps(cell.metrics, indent=2))
rows.extend(unsupported_rows)
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(rows, indent=2))
table_path = args.output_dir / "metrics.md"
table = _markdown_table(rows)
table_path.write_text(table + "\n")
html_path = _write_html_grid(rows, args.output_dir)
print("\n" + table)
print(f"\nWrote {metrics_path}, {table_path}, and {html_path}")
if failures:
raise SystemExit("SSIM regression gate failed:\n " + "\n ".join(failures))
def _load_lpips() -> object | None:
"""Return an LPIPS model, or ``None`` if the optional dep is unavailable."""
try:
import lpips # noqa: PLC0415 - optional dependency.
except ImportError:
print("LPIPS requested but the `lpips` package is not installed; skipping (install `.[eval]`).")
return None
return lpips.LPIPS(net="alex")
def _lpips_between(lpips_fn, reference_video: Path, candidate_video: Path) -> float | None:
if lpips_fn is None or not reference_video.exists() or not candidate_video.exists():
return None
import torch
ref = _read_video_frames(reference_video)
cand = _read_video_frames(candidate_video)
if ref is None or cand is None or ref.shape != cand.shape:
return None
# LPIPS expects NCHW in [-1, 1].
ref_t = torch.from_numpy(ref).permute(0, 3, 1, 2).float() / 127.5 - 1.0
cand_t = torch.from_numpy(cand).permute(0, 3, 1, 2).float() / 127.5 - 1.0
with torch.no_grad():
scores = lpips_fn(ref_t, cand_t)
return float(scores.mean().item())
def _read_video_frames(path: Path) -> np.ndarray | None:
try:
import cv2
except ImportError:
return None
cap = cv2.VideoCapture(str(path))
frames = []
try:
while True:
ok, frame_bgr = cap.read()
if not ok:
break
frames.append(cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB))
finally:
cap.release()
if not frames:
return None
return np.stack(frames, axis=0)
if __name__ == "__main__":
main()
@@ -62,14 +62,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
hidden_size: int = 5120
intermediate_size: int = 25600
num_hidden_layers: int = 64
# H3 conditions on one intermediate hidden state and reads nothing above it,
# so the remaining layers are built, weight-loaded and then discarded: 14
# layers, 13.7 GB in bf16. Building exactly this many leaves that hidden
# state bit-identical, because the tuple records each layer's *input*, so
# entry N is the output of layer N-1. Set to None to keep the full stack.
# Must equal MINIMAX_H3_TEXT_ENCODER_LAYER in
# fastvideo/pipelines/basic/minimax_h3/packing.py; a test pins them together
# rather than importing across the models -> pipelines boundary.
output_hidden_state_index: int = 50
num_hidden_layers_override: int | None = 50
num_attention_heads: int = 64
num_key_value_heads: int = 8
@@ -116,7 +109,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
vision_initializer_range: float = 0.02
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
output_hidden_states: bool = True
output_hidden_states: bool = False
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
_fsdp_shard_conditions: list = field(default_factory=lambda: [
_is_language_transformer_layer,
@@ -127,15 +120,16 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
])
def __post_init__(self) -> None:
# Runs both at construction and after ``update_model_arch`` merges the
# checkpoint's config.json, so it also guards config-file overrides. A
# non-positive override would build no decoder layers at all, and a
# negative one would additionally make the surplus-key filter drop
# every ``language_model.layers.*`` checkpoint key, so the conditioner
# would "load" with no transformer stack and only fail at generation.
if self.num_hidden_layers_override is not None and self.num_hidden_layers_override < 1:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be a positive layer count "
f"or None for the full stack; got {self.num_hidden_layers_override}.")
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
if self.num_hidden_layers_override is not None:
if self.num_hidden_layers_override <= 0:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
if self.num_hidden_layers_override < self.output_hidden_state_index:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
f"hidden_states[{self.output_hidden_state_index}], got "
f"{self.num_hidden_layers_override}.")
rope_scaling = dict(self.rope_scaling or {})
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
+43 -18
View File
@@ -808,14 +808,19 @@ class VideoGenerator:
latent_batch_size = _infer_latent_batch_size(batch)
is_latent_output = fastvideo_args.output_type == "latent"
needs_frame_output = batch.return_frames or (batch.save_video and not is_latent_output)
needs_samples_buffer = batch.return_frames or needs_frame_output
# When ``output_type == "latent"`` the forward output has latent
# shape (e.g. ``[B, C_latent, T_latent, H_latent, W_latent]``)
# rather than the pre-allocation's pixel shape. Skip the pinned
# ~50 MB buffer entirely. Also skip it for metadata-only calls;
# neither the result nor save path will consume the decoded tensor.
# A populated ``samples`` has exactly one consumer — the result
# dict (``"samples": samples if batch.return_frames else None``).
# Post-decode frame building reads ``output_batch.output``
# directly (the GPU ``vid_u8`` path), not ``samples``. So when
# ``return_frames=False`` the pinned fp32 alloc + D->H copy are
# dead weight — the CLI generate flow (``save_video=True``,
# ``return_frames=False``) hits this on every call.
# ``output_type == "latent"`` keeps its existing branch (shape
# mismatch falls through to ``.cpu()`` below) for callers that
# *do* ask for the latent samples via ``return_frames=True``.
# ``skip_pixel_prealloc`` also gates the slow-path warning.
skip_pixel_prealloc = is_latent_output or not needs_samples_buffer
needs_samples_out = batch.return_frames
skip_pixel_prealloc = is_latent_output or not needs_samples_out
if skip_pixel_prealloc:
samples = torch.empty(0, device='cpu')
else:
@@ -835,9 +840,11 @@ class VideoGenerator:
"This usually means the executor/pipeline failed earlier.")
audio_only = bool(output_batch.extra.get("audio_only"))
if not needs_samples_buffer or (audio_only and not batch.return_frames):
# Metadata-only/audio-only request: keep the empty placeholder and
# avoid the decoded tensor D->H copy.
if not needs_samples_out:
# Nothing downstream reads ``samples`` (the result dict
# returns None when ``return_frames=False``); keep the empty
# placeholder allocated above and skip the fp32 D->H copy
# entirely.
pass
elif audio_only:
# Audio-only return-frames requests expose the small placeholder
@@ -869,8 +876,13 @@ class VideoGenerator:
# `GenerationResult.size` describes the produced media, not only the
# base-stage request. Refiner pipelines can change the final pixel
# dimensions, so derive this result metadata from the decoded output.
# Read the geometry from `output_batch.output` (a shape-only access,
# no D->H copy): when `return_frames=False` the `samples` mirror
# stays an empty placeholder and no longer carries the decoded
# shape. Metadata-only calls keep the request fallback and never
# inspect the (possibly dropped) worker output.
output_size = _resolve_output_size(
samples,
output_batch.output if needs_frame_output else samples,
(target_height, target_width, batch.num_frames),
pixel_output=not is_latent_output and not audio_only,
)
@@ -882,13 +894,26 @@ class VideoGenerator:
elif not needs_frame_output:
frames = None
else:
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.permute(1, 2, 0).squeeze(-1)
x = (x * 255).to(torch.uint8)
frames.append(x.contiguous().cpu().numpy())
# Quantize on the source device (typically CUDA) BEFORE the
# device->host copy. `samples` above is just the pinned-CPU
# mirror of `output_batch.output` (`samples.copy_(output)` or
# `output.cpu()`) with no intervening preprocessing, so reading
# `output_batch.output` here is the same data. The old path
# paid a full fp32 video D->H copy (which scales with
# resolution x frames x batch) and then a single-threaded
# per-frame CPU *255/cast loop. Casting to uint8 on-device
# makes the transfer 4x smaller, ships it in a single copy,
# and moves the elementwise work onto the GPU. clamp_() also
# fixes a latent overflow bug: VAE output slightly outside
# [0, 1] wrapped mod 256 in the old unclamped cast.
# (Equivalence is SSIM-gated, not bit-exact: float->uint8
# differs <=1 LSB CPU vs GPU.)
src = output_batch.output
vid_u8 = (src * 255).clamp_(0, 255).to(torch.uint8)
vid_u8 = rearrange(vid_u8, "b c t h w -> t b c h w").cpu()
frames = [
torchvision.utils.make_grid(x, nrow=6).permute(1, 2, 0).squeeze(-1).contiguous().numpy() for x in vid_u8
]
postprocess_time = time.perf_counter() - postprocess_start
logger.info("PostDecodeFrameProcessStage completed in %.3f s", postprocess_time)
if logging_info is not None:
+29
View File
@@ -21,12 +21,17 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_MINIMAX_H3_FUSIONS: str = ""
FASTVIDEO_VAE_PARALLEL_DECODE: bool = False
FASTVIDEO_VAE_PARALLEL_ENCODE: bool = False
FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_NVTX_PROFILE: bool = False
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
@@ -217,10 +222,34 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# If set (=1), MiniMax-H3 VAE decode (and, with the ENCODE variant,
# reference-video encode) round-robins its temporal chunks across the
# sequence-parallel ranks instead of running serially on the output rank.
# Folded into FastVideoArgs.vae_parallel_decode / vae_parallel_encode at
# construction (parse-once). The STRATEGY variant picks the chunk
# transport collective: "gather" (default) or "all_gather".
"FASTVIDEO_VAE_PARALLEL_DECODE":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE", "0") != "0",
"FASTVIDEO_VAE_PARALLEL_ENCODE":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_ENCODE", "0") != "0",
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY", None),
# Opt-in MiniMax-H3 inference-only Triton fusions adapted from the
# NVlabs/Sana Sol-Engine implementation. Accepts `all`, `1`, or a
# comma-separated subset of `modulate,qknorm_rope,swiglu`. An empty value
# (the default), `0`, or `none` keeps the eager implementation.
"FASTVIDEO_MINIMAX_H3_FUSIONS":
lambda: os.getenv("FASTVIDEO_MINIMAX_H3_FUSIONS", ""),
# Use dedicated multiprocess context for workers.
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
# Emit lightweight NVTX ranges for external profilers such as Nsight Systems.
"FASTVIDEO_NVTX_PROFILE":
lambda: os.getenv("FASTVIDEO_NVTX_PROFILE", "0") != "0",
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
"FASTVIDEO_TORCH_PROFILER_DIR":
+51
View File
@@ -146,6 +146,19 @@ class FastVideoArgs:
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
# Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the
# video VAE's temporal chunks (decode) and clips (reference encode) are
# round-robined across the sequence-parallel ranks and reassembled
# bit-exactly on the group's first rank instead of running serially on
# one rank while the others idle. ``__post_init__`` folds the
# FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars
# into these fields (parse-once, like attention_backend), and
# FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport
# collective ("gather" or "all_gather").
vae_parallel_decode: bool = False
vae_parallel_encode: bool = False
vae_parallel_decode_strategy: str | None = None
# Compilation
# ``enable_torch_compile`` covers the DiT path (transformer,
# transformer_2, and the LTX-2 stage-2 transformer_refine).
@@ -169,6 +182,7 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
VSA_tile_size: int = 256 # VSA-H3 tile size (256 or 64); 64 = native Triton path
# V-MoBA parameters
moba_config_path: str | None = None
@@ -286,8 +300,27 @@ class FastVideoArgs:
env_backend = envs.FASTVIDEO_ATTENTION_BACKEND
if env_backend is not None and backend_name_to_enum(env_backend) is not None:
self.attention_backend = env_backend
self._fold_vae_parallel_env()
self.check_fastvideo_args()
def _fold_vae_parallel_env(self) -> None:
"""Parse-once adapters for the sequence-parallel VAE env vars."""
import fastvideo.envs as envs
# Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES /
# DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args
# never imports model modules; a unit test pins the two in sync).
strategies = ("gather", "all_gather")
if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE:
self.vae_parallel_decode = True
if not self.vae_parallel_encode and envs.FASTVIDEO_VAE_PARALLEL_ENCODE:
self.vae_parallel_encode = True
if self.vae_parallel_decode_strategy is None:
self.vae_parallel_decode_strategy = envs.FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY or "gather"
if self.vae_parallel_decode_strategy not in strategies:
raise ValueError(f"vae_parallel_decode_strategy must be one of {strategies}, "
f"got {self.vae_parallel_decode_strategy!r}.")
def _apply_transformer_quant(self) -> None:
"""Pin the typed ``transformer_quant`` instance onto ``dit_config``.
@@ -631,6 +664,18 @@ class FastVideoArgs:
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
"Should be enabled in almost all cases",
)
parser.add_argument(
"--vae-parallel-decode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 VAE decode chunks across the SP ranks "
"and reassemble bit-exactly on the output rank (default: serial decode on the output rank)",
)
parser.add_argument(
"--vae-parallel-encode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 reference-video VAE encode clips across "
"the SP ranks; every rank keeps the identical full encoding (default: serial encode on every rank)",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -644,6 +689,12 @@ class FastVideoArgs:
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
parser.add_argument(
"--VSA-tile-size",
type=int,
default=FastVideoArgs.VSA_tile_size,
help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path",
)
# Master port for distributed training/inference
parser.add_argument(
+4 -1
View File
@@ -114,7 +114,10 @@ def _info(logger: Logger,
is_local_main_process = local_rank == 0
if (main_process_only and is_main_process) or (local_main_process_only and is_local_main_process):
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
# Honor an explicit stacklevel (info_once routes through here with
# stacklevel already set) instead of passing the keyword twice.
stacklevel = kwargs.pop("stacklevel", 2)
logger.log(logging.INFO, msg, *args, stacklevel=stacklevel, **kwargs)
global _warned_local_main_process, _warned_main_process
+115
View File
@@ -0,0 +1,115 @@
# SPDX-License-Identifier: Apache-2.0
"""Experimental Apple MLX runtime helpers.
This package is intentionally small for now. It exists to grow the Apple-native
FastWan path in measurable steps: shape planning, primitive benchmarks, then
Wan block parity, then full DiT/runtime support.
"""
from fastvideo.mlx_runtime.fastwan import (
FastWanShape,
MLXQuantizationSpec,
MLXWanDiT,
MLXWanTransformerBlock,
UnsupportedMLXQuantizationError,
ensure_quantization_supported,
fastwan_shape,
fastwan_shape_from_config,
mlx_dit_from_diffusers_safetensors,
mlx_block_weights_from_torch,
mlx_block_weights_from_diffusers_safetensors,
quantization_support_error,
)
from fastvideo.mlx_runtime.checkpoint import (
load_mlx_dit_checkpoint,
save_mlx_dit_checkpoint,
)
from fastvideo.mlx_runtime.memory import (
AppliedMemoryLimits,
add_memory_limit_args,
apply_memory_limits,
gib_to_bytes,
)
from fastvideo.mlx_runtime.refine import (
DEFAULT_REFINE_SIGMA,
RefinePlan,
TwoPassResult,
default_refine_timesteps,
plan_refine_resolutions,
prepare_refine_latents,
refine_sigma_from_schedule,
run_dmd_loop,
run_two_pass_dmd,
upsample_latents_spatial,
)
from fastvideo.mlx_runtime.frame_upsample import (
DEFAULT_PIXEL_UPSAMPLE_MODE,
PIXEL_UPSAMPLE_MODES,
unsharp,
upsample_frame,
upsample_frames,
)
from fastvideo.mlx_runtime.fast_spatial import (
DEFAULT_FAST_SPATIAL_SHARPEN,
FastSpatialPlan,
apply_fast_spatial_upsample,
plan_fast_spatial,
resolve_spatial_mode,
)
from fastvideo.mlx_runtime.prompt_enhance import (
DEFAULT_ENHANCE_SYSTEM_PROMPT,
DEFAULT_MLX_LM_MODEL,
EnhanceResult,
enhance_prompt,
enhance_prompt_template,
enhance_result_as_metrics,
load_or_enhance_prompt,
)
__all__ = [
"AppliedMemoryLimits",
"DEFAULT_ENHANCE_SYSTEM_PROMPT",
"DEFAULT_MLX_LM_MODEL",
"DEFAULT_FAST_SPATIAL_SHARPEN",
"DEFAULT_PIXEL_UPSAMPLE_MODE",
"DEFAULT_REFINE_SIGMA",
"EnhanceResult",
"FastSpatialPlan",
"FastWanShape",
"MLXQuantizationSpec",
"MLXWanDiT",
"MLXWanTransformerBlock",
"RefinePlan",
"TwoPassResult",
"UnsupportedMLXQuantizationError",
"add_memory_limit_args",
"apply_fast_spatial_upsample",
"apply_memory_limits",
"enhance_prompt",
"enhance_prompt_template",
"enhance_result_as_metrics",
"ensure_quantization_supported",
"fastwan_shape",
"fastwan_shape_from_config",
"gib_to_bytes",
"load_mlx_dit_checkpoint",
"load_or_enhance_prompt",
"mlx_dit_from_diffusers_safetensors",
"mlx_block_weights_from_diffusers_safetensors",
"mlx_block_weights_from_torch",
"PIXEL_UPSAMPLE_MODES",
"default_refine_timesteps",
"plan_fast_spatial",
"plan_refine_resolutions",
"prepare_refine_latents",
"quantization_support_error",
"refine_sigma_from_schedule",
"resolve_spatial_mode",
"run_dmd_loop",
"run_two_pass_dmd",
"save_mlx_dit_checkpoint",
"unsharp",
"upsample_frame",
"upsample_frames",
"upsample_latents_spatial",
]
+273
View File
@@ -0,0 +1,273 @@
# SPDX-License-Identifier: Apache-2.0
"""Pre-quantized MLX checkpoint save/load for the FastWan DiT.
Loading the Diffusers fp32/fp16 checkpoint and quantizing at startup costs
both download size and load time on every run. This module persists an already
cast (and optionally already quantized) ``MLXWanDiT`` so 16 GB users download
and load roughly half the bytes and skip requantization entirely:
dit = mlx_dit_from_diffusers_safetensors(ckpt, cfg, quantization="int8")
save_mlx_dit_checkpoint(dit, "FastWan2.1-T2V-1.3B-mlx-int8")
...
dit = load_mlx_dit_checkpoint("FastWan2.1-T2V-1.3B-mlx-int8")
Format (one directory):
- ``mlx_dit.safetensors`` — every array, saved with ``mx.save_safetensors``.
Plain weights keep their key; a quantized weight ``K`` is stored as the
packed ``K`` plus ``K.scales`` (and ``K.biases`` for affine modes).
- ``mlx_dit.json`` — format version, the model config, the quantization spec,
and which keys are quantized, so the loader can rebuild ``QuantizedMatrix``
objects without guessing.
"""
from __future__ import annotations
import json
import shutil
import tempfile
from pathlib import Path
from typing import Any
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.fastwan import (
MLXQuantizationSpec,
MLXWanDiT,
MLXWanTransformerBlock,
QuantizedMatrix,
ensure_quantization_supported,
)
logger = init_logger(__name__)
FORMAT_VERSION = 1
WEIGHTS_FILENAME = "mlx_dit.safetensors"
MANIFEST_FILENAME = "mlx_dit.json"
_BLOCK_PREFIX = "blocks"
_DTYPE_TO_NAME = {"float16": "fp16", "bfloat16": "bf16", "float32": "fp32"}
def _dtype_name(dtype) -> str:
"""Return the manifest name for a supported MLX data type.
Parameters:
dtype: The MLX data type to convert.
Returns:
str: The manifest name corresponding to the data type.
Raises:
ValueError: If the data type is not supported for checkpointing.
"""
import mlx.core as mx
for raw, name in _DTYPE_TO_NAME.items():
if dtype == getattr(mx, raw):
return name
raise ValueError(f"Unsupported MLX dtype for checkpointing: {dtype}")
def _name_to_dtype(name: str):
"""Convert a manifest dtype name to its corresponding MLX dtype.
Parameters:
name (str): Manifest name, such as ``"fp16"``, ``"bf16"``, or ``"fp32"``.
Returns:
The corresponding MLX dtype.
"""
import mlx.core as mx
return {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[name]
def _flatten_weights(dit: MLXWanDiT) -> dict[str, Any]:
"""Combine model and transformer-block weights into a single flattened mapping.
Parameters:
dit (MLXWanDiT): Model whose weights should be flattened.
Returns:
dict[str, Any]: Mapping containing top-level weights and indexed transformer-block weights.
"""
flat: dict[str, Any] = dict(dit.weights)
for index, block in enumerate(dit.blocks):
for name, value in block.weights.items():
flat[f"{_BLOCK_PREFIX}.{index}.{name}"] = value
return flat
def save_mlx_dit_checkpoint(dit: MLXWanDiT, checkpoint_dir: str | Path) -> Path:
"""Save a plain or quantized MLX Wan DiT checkpoint to a directory.
Parameters:
dit (MLXWanDiT): Model whose weights and configuration will be saved.
checkpoint_dir (str | Path): Destination directory for the checkpoint.
Returns:
Path: Path to the checkpoint directory.
"""
import mlx.core as mx
checkpoint_dir = Path(checkpoint_dir)
arrays: dict[str, Any] = {}
quantized: dict[str, dict[str, Any]] = {}
spec: MLXQuantizationSpec | None = None
for key, value in _flatten_weights(dit).items():
if isinstance(value, QuantizedMatrix):
if spec is not None and value.spec != spec:
raise ValueError(f"Mixed quantization specs in one checkpoint ({spec} vs {value.spec} at '{key}') "
"are not supported.")
spec = value.spec
arrays[key] = value.weight
arrays[f"{key}.scales"] = value.scales
if value.biases is not None:
arrays[f"{key}.biases"] = value.biases
quantized[key] = {
"dequantized_dtype": _dtype_name(value.dequantized_dtype),
"has_biases": value.biases is not None,
}
else:
arrays[key] = value
manifest = {
"format_version": FORMAT_VERSION,
"config": dit.config,
"num_blocks": len(dit.blocks),
"quantization": None if spec is None else {
"mode": spec.mode,
"bits": spec.bits,
"group_size": spec.group_size,
},
"quantized_keys": quantized,
}
manifest_json = json.dumps(manifest, indent=2)
checkpoint_dir.parent.mkdir(parents=True, exist_ok=True)
staging_dir = Path(tempfile.mkdtemp(dir=checkpoint_dir.parent, prefix=f".{checkpoint_dir.name}.staging-"))
backup_root: Path | None = None
try:
staged_weights = staging_dir / WEIGHTS_FILENAME
staged_manifest = staging_dir / MANIFEST_FILENAME
mx.save_safetensors(str(staged_weights), arrays)
staged_manifest.write_text(manifest_json)
if checkpoint_dir.exists():
backup_root = Path(tempfile.mkdtemp(dir=checkpoint_dir.parent, prefix=f".{checkpoint_dir.name}.backup-"))
try:
checkpoint_dir.replace(backup_root / checkpoint_dir.name)
except Exception:
shutil.rmtree(backup_root, ignore_errors=True)
raise
try:
staging_dir.replace(checkpoint_dir)
except Exception:
if backup_root is not None:
(backup_root / checkpoint_dir.name).replace(checkpoint_dir)
shutil.rmtree(backup_root, ignore_errors=True)
raise
if backup_root is not None:
shutil.rmtree(backup_root, ignore_errors=True)
finally:
shutil.rmtree(staging_dir, ignore_errors=True)
logger.info("Saved MLX DiT checkpoint (%d arrays, quantization=%s) to %s", len(arrays),
spec.label if spec else "none", checkpoint_dir)
return checkpoint_dir
def load_mlx_dit_checkpoint(checkpoint_dir: str | Path, *, compile: bool = False) -> MLXWanDiT:
"""
Reconstruct an MLXWanDiT model from a versioned checkpoint.
Parameters:
checkpoint_dir (str | Path): Directory containing the checkpoint manifest and weights.
compile (bool): Whether to configure the reconstructed model for compilation.
Returns:
MLXWanDiT: The reconstructed model.
Raises:
FileNotFoundError: If the checkpoint manifest or weights file is missing.
ValueError: If the checkpoint format is unsupported or block weights are incomplete.
"""
import mlx.core as mx
checkpoint_dir = Path(checkpoint_dir)
manifest_path = checkpoint_dir / MANIFEST_FILENAME
weights_path = checkpoint_dir / WEIGHTS_FILENAME
if not manifest_path.exists() or not weights_path.exists():
raise FileNotFoundError(f"Not an MLX DiT checkpoint directory: {checkpoint_dir} "
f"(expected {MANIFEST_FILENAME} and {WEIGHTS_FILENAME}).")
manifest = json.loads(manifest_path.read_text())
version = manifest.get("format_version")
if version != FORMAT_VERSION:
raise ValueError(f"MLX DiT checkpoint {checkpoint_dir} has format_version={version}; "
f"this FastVideo build reads version {FORMAT_VERSION}. Re-export the checkpoint.")
spec = None
if manifest["quantization"] is not None:
spec = MLXQuantizationSpec(**manifest["quantization"])
# The packed layout of mx.quantize output is mode-specific, so a build
# that cannot run the mode cannot use these arrays at all.
ensure_quantization_supported(spec)
arrays = mx.load(str(weights_path))
quantized_keys: dict[str, dict[str, Any]] = manifest["quantized_keys"]
def rebuild(key: str):
"""
Reconstructs a weight array or quantized matrix from checkpoint data.
Parameters:
key (str): The weight key to rebuild.
Returns:
The stored array for an unquantized weight or a reconstructed quantized matrix.
"""
if key not in quantized_keys:
return arrays[key]
info = quantized_keys[key]
assert spec is not None, f"Quantized key '{key}' in a checkpoint without a quantization spec"
return QuantizedMatrix(
weight=arrays[key],
scales=arrays[f"{key}.scales"],
biases=arrays[f"{key}.biases"] if info["has_biases"] else None,
spec=spec,
dequantized_dtype=_name_to_dtype(info["dequantized_dtype"]),
)
config = manifest["config"]
block_keys: dict[int, list[str]] = {}
top_level_keys: list[str] = []
for key in arrays:
if key.endswith(".scales") or key.endswith(".biases"):
continue
if key.startswith(f"{_BLOCK_PREFIX}."):
index_str, _, _ = key[len(_BLOCK_PREFIX) + 1:].partition(".")
block_keys.setdefault(int(index_str), []).append(key)
else:
top_level_keys.append(key)
weights = {key: rebuild(key) for key in top_level_keys}
num_blocks = int(manifest["num_blocks"])
if sorted(block_keys) != list(range(num_blocks)):
raise ValueError(f"MLX DiT checkpoint {checkpoint_dir} is missing block weights: "
f"manifest says {num_blocks} blocks, found indices {sorted(block_keys)}.")
inner_dim = int(config["num_attention_heads"]) * int(config["attention_head_dim"])
blocks = []
for index in range(num_blocks):
prefix = f"{_BLOCK_PREFIX}.{index}."
block_weights = {key[len(prefix):]: rebuild(key) for key in block_keys[index]}
blocks.append(
MLXWanTransformerBlock(
block_weights,
dim=inner_dim,
ffn_dim=int(config["ffn_dim"]),
num_heads=int(config["num_attention_heads"]),
eps=float(config["eps"]),
))
return MLXWanDiT(weights, blocks, config, compile=compile)
+228
View File
@@ -0,0 +1,228 @@
# SPDX-License-Identifier: Apache-2.0
"""Spatial fast mode for the MLX Wan runtime (RIFE's spatial twin).
RIFE ``--fast`` cuts *frames* (temporal). This module cuts *pixels*
(spatial): denoise at ``target // scale``, decode at that size, then
resample the decoded frames up to the target. No second denoise pass —
that is ``--refine`` (quality). The two compose:
* ``--fast-spatial`` alone → speed (≈ scale² fewer tokens)
* ``--refine`` alone → quality two-pass (H3 / LTX-2)
* ``--fast`` + ``--refine`` → fewer frames at base res, full-res refine
* ``--fast`` + ``--fast-spatial`` → fewer frames *and* fewer pixels
The upsample runs in **pixel** space, after the VAE decode. It used to run
in latent space (bilinear over the latent H/W plane, sharing the refine
hand-off primitive) and that is what made spatial fast mode incoherent: an
interpolated Wan latent is off the decoder's manifold, so decode returned
the right silhouette under a smeared veil. ``--refine`` can get away with
the latent-space upsample because a second DMD pass re-denoises the result;
spatial fast mode hands the latent straight to the decoder, so it cannot.
See :mod:`fastvideo.mlx_runtime.frame_upsample` for the full rationale.
MetalFX is intentionally not used: it needs game-engine motion vectors
and depth that diffusion output lacks.
"""
from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass
import numpy as np
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.frame_upsample import (
DEFAULT_PIXEL_UPSAMPLE_MODE,
PIXEL_UPSAMPLE_MODES,
upsample_frames,
)
from fastvideo.mlx_runtime.refine import RefinePlan, plan_refine_resolutions
logger = init_logger(__name__)
# Resampling from a smaller decode loses high-frequency detail the same way
# RIFE's flow warp does, so spatial fast mode borrows ``--fast``'s remedy: a
# light unsharp pass. 0.4 recovers perceived crispness on Wan2.1 output at 2x
# without the halos that show up by ~0.8.
DEFAULT_FAST_SPATIAL_SHARPEN = 0.4
@dataclass(frozen=True)
class FastSpatialPlan:
"""Resolved geometry for a spatial-fast (upsample-only) run."""
plan: RefinePlan
upsample_mode: str
sharpen: float = DEFAULT_FAST_SPATIAL_SHARPEN
@property
def enabled(self) -> bool:
"""
Determine whether spatial scaling is enabled.
Returns:
`true` if the spatial scale is greater than one, `false` otherwise.
"""
return self.plan.spatial_scale > 1
@property
def scale(self) -> int:
"""Provides the configured spatial scaling factor.
Returns:
int: The spatial scaling factor.
"""
return self.plan.spatial_scale
@property
def target_height(self) -> int:
"""
Return the target output height for the spatial plan.
Returns:
int: Target output height in pixels.
"""
return self.plan.target_height
@property
def target_width(self) -> int:
"""Return the target image width in pixels.
Returns:
int: The target image width.
"""
return self.plan.target_width
@property
def stage1_height(self) -> int:
"""
Provide the stage-one latent height used for reduced-resolution processing.
Returns:
int: The stage-one latent height.
"""
return self.plan.stage1_height
@property
def stage1_width(self) -> int:
"""Get the stage-one latent width.
Returns:
int: The stage-one latent width.
"""
return self.plan.stage1_width
def plan_fast_spatial(
*,
height: int,
width: int,
num_frames: int,
spatial_scale: int = 2,
vae_spatial_compression: int = 8,
vae_temporal_compression: int = 4,
patch_size: tuple[int, int, int] = (1, 2, 2),
upsample_mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
sharpen: float = DEFAULT_FAST_SPATIAL_SHARPEN,
enabled: bool = True,
) -> FastSpatialPlan:
"""
Build a plan for reduced-resolution denoising followed by pixel-space upsampling.
Parameters:
upsample_mode (str): Pixel interpolation kernel, one of
:data:`~fastvideo.mlx_runtime.frame_upsample.PIXEL_UPSAMPLE_MODES`.
sharpen (float): Unsharp strength applied after the resize.
Returns:
FastSpatialPlan: The validated spatial-fast processing plan.
Raises:
ValueError: If the upsample mode is unsupported or ``sharpen`` is negative.
"""
if upsample_mode not in PIXEL_UPSAMPLE_MODES:
raise ValueError(f"Unsupported upsample mode: {upsample_mode!r} "
f"(expected one of {', '.join(PIXEL_UPSAMPLE_MODES)})")
if sharpen < 0.0:
raise ValueError(f"sharpen must be >= 0, got {sharpen}")
plan = plan_refine_resolutions(
height=height,
width=width,
num_frames=num_frames,
spatial_scale=spatial_scale,
vae_spatial_compression=vae_spatial_compression,
vae_temporal_compression=vae_temporal_compression,
patch_size=patch_size,
enabled=enabled,
mode_label="fast-spatial",
)
if plan.spatial_scale > 1:
logger.info(
"[MLX fast-spatial] denoise+decode %dx%d → upsample %dx to %dx%d (%s, sharpen=%.2f)",
plan.stage1_width,
plan.stage1_height,
plan.spatial_scale,
plan.target_width,
plan.target_height,
upsample_mode,
sharpen,
)
return FastSpatialPlan(plan=plan, upsample_mode=upsample_mode, sharpen=sharpen)
def apply_fast_spatial_upsample(
frames: Iterable[np.ndarray],
spatial: FastSpatialPlan,
) -> list[np.ndarray]:
"""Resample decoded stage-1 frames up to the target resolution.
This runs on decoded RGB frames, *not* on latents: see the module
docstring for why the latent-space version produced a blurred veil.
Parameters:
frames (Iterable[np.ndarray]): Decoded HxWx3 uint8 RGB frames, produced
by decoding at the stage-one resolution.
spatial (FastSpatialPlan): Plan defining the target size, interpolation
kernel, and unsharp strength.
Returns:
list[np.ndarray]: Frames at the target resolution. When spatial scaling
is disabled the frames are returned unchanged, as a list.
"""
if not spatial.enabled:
return list(frames)
return upsample_frames(
frames,
width=spatial.target_width,
height=spatial.target_height,
mode=spatial.upsample_mode,
sharpen=spatial.sharpen,
)
def resolve_spatial_mode(
*,
refine: bool,
fast_spatial: bool,
) -> str:
"""Select the active spatial processing mode, with refinement taking precedence.
Returns:
str: ``"refine"`` when refinement is enabled, ``"fast_spatial"`` when
spatial-fast processing is enabled, or ``"off"`` otherwise.
"""
if refine:
return "refine"
if fast_spatial:
return "fast_spatial"
return "off"
__all__ = [
"DEFAULT_FAST_SPATIAL_SHARPEN",
"FastSpatialPlan",
"apply_fast_spatial_upsample",
"plan_fast_spatial",
"resolve_spatial_mode",
]
+980
View File
@@ -0,0 +1,980 @@
# SPDX-License-Identifier: Apache-2.0
# mypy: disable-error-code=no-untyped-call
"""FastWan-oriented helpers for the experimental MLX runtime path."""
from __future__ import annotations
import json
import math
import os
import statistics
import time
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any
from collections.abc import Callable
from fastvideo.logger import init_logger
if TYPE_CHECKING:
import mlx.core as mx
import torch
logger = init_logger(__name__)
@dataclass(frozen=True)
class FastWanShape:
height: int
width: int
num_frames: int
latent_frames: int
latent_height: int
latent_width: int
patch_frames: int
patch_height: int
patch_width: int
tokens: int
hidden_size: int
num_heads: int
head_dim: int
class UnsupportedMLXQuantizationError(ValueError):
"""A quantization mode the installed MLX build cannot execute.
Raised by :func:`ensure_quantization_supported` before any model weights
are loaded, so callers (CLI flags, benchmark sweeps) can fail fast with an
actionable message -- or skip the mode -- instead of crashing deep inside
``mx.quantize`` mid-load.
"""
@dataclass(frozen=True)
class MLXQuantizationSpec:
"""MLX quantized-matmul configuration for DiT linear weights."""
mode: str
bits: int | None = None
group_size: int | None = None
@classmethod
def from_name(cls, name: str | None) -> MLXQuantizationSpec | None:
if name is None or name in {"", "none", "fp16", "fp32"}:
return None
if name == "int8":
return cls(mode="affine", bits=8, group_size=64)
if name == "int4":
return cls(mode="affine", bits=4, group_size=64)
if name == "mxfp8":
return cls(mode="mxfp8")
if name == "mxfp4":
return cls(mode="mxfp4")
if name == "nvfp4":
return cls(mode="nvfp4")
raise ValueError(f"Unsupported MLX quantization mode: {name}")
@property
def label(self) -> str:
if self.mode == "affine":
return f"int{self.bits}"
return self.mode
@dataclass(frozen=True)
class QuantizedMatrix:
weight: mx.array
scales: mx.array
biases: mx.array | None
spec: MLXQuantizationSpec
dequantized_dtype: mx.Dtype
def fastwan_shape(
*,
height: int,
width: int,
num_frames: int,
vae_temporal_compression: int = 4,
vae_spatial_compression: int = 8,
patch_size: tuple[int, int, int] = (1, 2, 2),
num_heads: int = 12,
head_dim: int = 128,
) -> FastWanShape:
"""Return the approximate DiT token shape for Wan/FastWan T2V inference."""
latent_frames = (num_frames - 1) // vae_temporal_compression + 1
latent_height = height // vae_spatial_compression
latent_width = width // vae_spatial_compression
patch_frames = latent_frames // patch_size[0]
patch_height = latent_height // patch_size[1]
patch_width = latent_width // patch_size[2]
tokens = patch_frames * patch_height * patch_width
return FastWanShape(
height=height,
width=width,
num_frames=num_frames,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
patch_frames=patch_frames,
patch_height=patch_height,
patch_width=patch_width,
tokens=tokens,
hidden_size=num_heads * head_dim,
num_heads=num_heads,
head_dim=head_dim,
)
def fastwan_shape_from_config(
config_path: str | Path,
*,
height: int,
width: int,
num_frames: int,
) -> FastWanShape:
config = json.loads(Path(config_path).read_text())
return fastwan_shape(
height=height,
width=width,
num_frames=num_frames,
patch_size=tuple(config["patch_size"]),
num_heads=int(config["num_attention_heads"]),
head_dim=int(config["attention_head_dim"]),
)
def replace_tokens(shape: FastWanShape, tokens: int) -> FastWanShape:
return FastWanShape(**{**shape.__dict__, "tokens": tokens})
def median_ms(samples: list[float]) -> float:
return statistics.median(samples) * 1000.0
def benchmark_mlx_attention(shape: FastWanShape, warmup: int, iters: int) -> float:
import mlx.core as mx
q = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
k = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
v = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
scale = shape.head_dim**-0.5
for _ in range(warmup):
y = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
mx.eval(y)
samples = []
for _ in range(iters):
start = time.perf_counter()
y = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
mx.eval(y)
samples.append(time.perf_counter() - start)
return median_ms(samples)
def benchmark_mlx_linear(shape: FastWanShape, warmup: int, iters: int) -> float:
import mlx.core as mx
x = mx.random.normal((shape.tokens, shape.hidden_size), dtype=mx.float16)
w = mx.random.normal((shape.hidden_size, shape.hidden_size), dtype=mx.float16)
b = mx.zeros((shape.hidden_size, ), dtype=mx.float16)
for _ in range(warmup):
y = x @ w + b
mx.eval(y)
samples = []
for _ in range(iters):
start = time.perf_counter()
y = x @ w + b
mx.eval(y)
samples.append(time.perf_counter() - start)
return median_ms(samples)
def benchmark_torch_mps_attention(shape: FastWanShape, warmup: int, iters: int) -> float | None:
try:
import torch
import torch.nn.functional as F
except ImportError:
return None
if not torch.backends.mps.is_available():
return None
device = torch.device("mps")
q = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
k = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
v = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
for _ in range(warmup):
y = F.scaled_dot_product_attention(q, k, v)
torch.mps.synchronize()
_ = y
samples = []
for _ in range(iters):
start = time.perf_counter()
y = F.scaled_dot_product_attention(q, k, v)
torch.mps.synchronize()
_ = y
samples.append(time.perf_counter() - start)
return median_ms(samples)
def torch_to_mx(tensor) -> mx.array:
import mlx.core as mx
return mx.array(tensor.detach().cpu().float().numpy())
def weight_dtype(weight):
if isinstance(weight, QuantizedMatrix):
return weight.dequantized_dtype
return weight.dtype
_QUANT_SUPPORT_CACHE: dict[tuple[str, int | None, int | None], str | None] = {}
def quantization_support_error(spec: MLXQuantizationSpec) -> str | None:
"""Probe whether the installed MLX build supports ``spec``.
Runs a tiny ``mx.quantize`` + ``mx.quantized_matmul`` with exactly the
arguments :func:`quantize_matrix` / :func:`linear` use, so the result
reflects the real runtime path. The affine (int8/int4) modes are stable
across MLX releases, but the ``mxfp8``/``mxfp4``/``nvfp4`` mode strings
require newer MLX builds and raise otherwise. Returns ``None`` when the
mode works, else the underlying error message. Cached per spec.
"""
key = (spec.mode, spec.bits, spec.group_size)
if key not in _QUANT_SUPPORT_CACHE:
import mlx.core as mx
try:
probe_dim = max(spec.group_size or 0, 64)
weight = mx.zeros((probe_dim, probe_dim), dtype=mx.float16)
quantized = quantize_matrix(weight, spec)
y = linear(mx.zeros((1, probe_dim), dtype=mx.float16), quantized)
mx.eval(y)
_QUANT_SUPPORT_CACHE[key] = None
except Exception as exc: # noqa: BLE001 - MLX raises varied error types per backend/version.
_QUANT_SUPPORT_CACHE[key] = f"{type(exc).__name__}: {exc}"
return _QUANT_SUPPORT_CACHE[key]
def ensure_quantization_supported(spec: MLXQuantizationSpec | None) -> None:
"""Raise :class:`UnsupportedMLXQuantizationError` if ``spec`` cannot run here."""
if spec is None:
return
error = quantization_support_error(spec)
if error is None:
return
import mlx.core as mx
mlx_version = getattr(mx, "__version__", "unknown")
raise UnsupportedMLXQuantizationError(f"MLX quantization mode '{spec.label}' is not supported by the installed mlx "
f"({mlx_version}): {error}. Upgrade mlx or pick a supported mode "
f"(int8 is currently the most reliable quality/memory target).")
def quantize_matrix(weight, spec: MLXQuantizationSpec | None):
if spec is None:
return weight
import mlx.core as mx
if len(weight.shape) < 2:
return weight
q = mx.quantize(weight, group_size=spec.group_size, bits=spec.bits, mode=spec.mode)
biases = q[2] if len(q) == 3 else None
eval_args = [q[0], q[1]]
if biases is not None:
eval_args.append(biases)
mx.eval(*eval_args)
return QuantizedMatrix(
weight=q[0],
scales=q[1],
biases=biases,
spec=spec,
dequantized_dtype=weight.dtype,
)
def linear(x, weight, bias=None):
import mlx.core as mx
if isinstance(weight, QuantizedMatrix):
y = mx.quantized_matmul(
x,
weight.weight,
weight.scales,
weight.biases,
transpose=True,
group_size=weight.spec.group_size,
bits=weight.spec.bits,
mode=weight.spec.mode,
).astype(x.dtype)
else:
y = x @ weight.T
if bias is not None:
y = y + bias
return y
def _use_fast_norm() -> bool:
"""Opt-in to MLX's fused ``mx.fast`` normalization kernels.
Off by default so the numerically-explicit reference path stays the
baseline. Set ``FASTVIDEO_MLX_FAST_NORM=1`` to route LayerNorm/RMSNorm
through single fused Metal kernels (fewer intermediates, less memory
traffic) and benchmark the speedup.
"""
import os
return os.environ.get("FASTVIDEO_MLX_FAST_NORM", "0") == "1"
def layer_norm(x, weight=None, bias=None, eps: float = 1e-6):
import mlx.core as mx
if _use_fast_norm():
# Compute in fp32 (matching the reference below) so downstream dtype
# and precision are identical across call sites.
w = weight.astype(mx.float32) if weight is not None else None
b = bias.astype(mx.float32) if bias is not None else None
return mx.fast.layer_norm(x.astype(mx.float32), w, b, eps)
x_float = x.astype(mx.float32)
mean = mx.mean(x_float, axis=-1, keepdims=True)
var = mx.mean(mx.square(x_float - mean), axis=-1, keepdims=True)
y = (x_float - mean) * mx.rsqrt(var + eps)
if weight is not None:
y = y * weight
if bias is not None:
y = y + bias
return y
def rms_norm(x, weight, eps: float = 1e-6):
import mlx.core as mx
if _use_fast_norm():
return mx.fast.rms_norm(x, weight, eps)
orig_dtype = x.dtype
x_float = x.astype(mx.float32)
variance = mx.mean(mx.square(x_float), axis=-1, keepdims=True)
y = x_float * mx.rsqrt(variance + eps)
return y.astype(orig_dtype) * weight
def apply_rotary_emb(x, cos, sin, *, is_neox_style: bool = False):
"""Apply FastVideo's rotary convention to MLX tensors.
Args:
x: [batch, seq, heads, head_dim]
cos/sin: [seq, head_dim] for Wan's full-dimension rotate-pair style,
or [seq, head_dim // 2] for traditional RoPE.
"""
import mlx.core as mx
head_size = x.shape[-1]
rope_dim = cos.shape[-1]
cos = cos[None, :, None, :]
sin = sin[None, :, None, :]
x_float = x.astype(mx.float32)
if rope_dim == head_size:
x_pairs = x_float.reshape(*x.shape[:-1], -1, 2)
x_real = x_pairs[..., 0]
x_imag = x_pairs[..., 1]
x_rotated = mx.stack([-x_imag, x_real], axis=-1).reshape(*x.shape)
return (x_float * cos + x_rotated * sin).astype(x.dtype)
if is_neox_style:
x1, x2 = mx.split(x_float, 2, axis=-1)
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
return mx.concatenate([o1, o2], axis=-1).astype(x.dtype)
x1 = x_float[..., ::2]
x2 = x_float[..., 1::2]
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
return mx.stack([o1, o2], axis=-1).reshape(*x.shape).astype(x.dtype)
_WINDOWED_ATTENTION_WARNED = False
def _warn_windowed_attention_once(window: int) -> None:
"""Warn that FASTVIDEO_MLX_WINDOW degrades output on a dense-trained DiT.
Sliding-window self-attention is fast (6.6x at a +-3-frame window on 1.3B)
but these checkpoints were trained with dense attention, and restricting it
at inference produces heavy colour-block noise: structural agreement with
the dense baseline drops to 0.25 at +-3 frames and 0.03 at +-5. Sparsity of
this kind is a training-time method. Kept as a research knob, but it should
never be on by accident.
"""
global _WINDOWED_ATTENTION_WARNED
if _WINDOWED_ATTENTION_WARNED:
return
_WINDOWED_ATTENTION_WARNED = True
logger.warning(
"FASTVIDEO_MLX_WINDOW=%d enables sliding-window self-attention. These "
"checkpoints are trained dense; expect severely degraded output. This is "
"a research knob, not a speed setting — use --fast-spatial for real "
"denoise savings.",
window,
)
def gelu_tanh(x):
"""tanh-approximate GELU, as used by Wan's FFN.
``mlx.nn.gelu_approx`` is the same tanh approximation behind a fused
kernel. On the 1.3B FFN shape (32760x8960) it is bit-identical to the
expanded expression below and 3.3x faster — 28.9ms -> 8.7ms per layer,
which is 0.6s per denoise step across 30 layers.
"""
import mlx.nn as nn
return nn.gelu_approx(x)
def silu(x):
import mlx.core as mx
return x * mx.sigmoid(x)
def timestep_embedding(t, dim: int, max_period: int = 10000):
import mlx.core as mx
half = dim // 2
freqs = mx.exp(-math.log(max_period) * mx.arange(0, half, dtype=mx.float32) / half)
args = t[:, None].astype(mx.float32) * freqs[None]
embedding = mx.concatenate([mx.cos(args), mx.sin(args)], axis=-1)
if dim % 2:
embedding = mx.concatenate([embedding, mx.zeros_like(embedding[:, :1])], axis=-1)
return embedding
def scale_residual(residual, x, gate):
return residual + x * gate
def scale_residual_layer_norm_scale_shift(residual, x, gate, shift, scale, weight=None, bias=None, eps: float = 1e-6):
if isinstance(gate, int):
assert gate == 1
residual_output = residual + x
else:
residual_output = residual + x * gate
normalized = layer_norm(residual_output, weight=weight, bias=bias, eps=eps)
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
class MLXWanT2VCrossAttention:
def __init__(self, weights: dict[str, mx.array], *, dim: int, num_heads: int, eps: float = 1e-6) -> None:
self.weights = weights
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
def __call__(self, x, context):
import mlx.core as mx
batch = x.shape[0]
q = linear(x, self.weights["attn2.to_q.weight"], self.weights.get("attn2.to_q.bias"))
q = rms_norm(q, self.weights["attn2.norm_q.weight"], eps=self.eps).reshape(batch, -1, self.num_heads,
self.head_dim)
if context.shape[1] == 0:
attended = mx.zeros_like(q)
else:
k = linear(context, self.weights["attn2.to_k.weight"], self.weights.get("attn2.to_k.bias"))
k = rms_norm(k, self.weights["attn2.norm_k.weight"],
eps=self.eps).reshape(batch, -1, self.num_heads, self.head_dim)
v = linear(context, self.weights["attn2.to_v.weight"],
self.weights.get("attn2.to_v.bias")).reshape(batch, -1, self.num_heads, self.head_dim)
attended = mx.fast.scaled_dot_product_attention(
q.transpose(0, 2, 1, 3),
k.transpose(0, 2, 1, 3),
v.transpose(0, 2, 1, 3),
scale=self.head_dim**-0.5,
).transpose(0, 2, 1, 3)
attended = attended.reshape(batch, -1, self.dim)
return linear(attended, self.weights["attn2.to_out.weight"], self.weights.get("attn2.to_out.bias"))
class MLXWanTransformerBlock:
"""Dense T2V Wan transformer block for the experimental MLX runtime.
This mirrors the non-VSA PyTorch block for single-process dense attention.
Rotary embeddings and sequence-parallel paths are intentionally left out of
this first parity target.
"""
def __init__(self, weights: dict[str, mx.array], *, dim: int, ffn_dim: int, num_heads: int, eps: float = 1e-6):
self.weights = weights
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
self.attn2 = MLXWanT2VCrossAttention(weights, dim=dim, num_heads=num_heads, eps=eps)
def __call__(self, hidden_states, encoder_hidden_states, temb, freqs_cis=None):
import mlx.core as mx
orig_dtype = hidden_states.dtype
e = self.weights["scale_shift_table"] + temb.astype(mx.float32)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = mx.split(e, 6, axis=1)
norm_hidden_states = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm_hidden_states = (norm_hidden_states * (1.0 + scale_msa) + shift_msa).astype(orig_dtype)
query = linear(norm_hidden_states, self.weights["to_q.weight"], self.weights.get("to_q.bias"))
key = linear(norm_hidden_states, self.weights["to_k.weight"], self.weights.get("to_k.bias"))
value = linear(norm_hidden_states, self.weights["to_v.weight"], self.weights.get("to_v.bias"))
query = rms_norm(query, self.weights["norm_q.weight"],
eps=self.eps).reshape(hidden_states.shape[0], -1, self.num_heads, self.head_dim)
key = rms_norm(key, self.weights["norm_k.weight"], eps=self.eps).reshape(hidden_states.shape[0], -1,
self.num_heads, self.head_dim)
value = value.reshape(hidden_states.shape[0], -1, self.num_heads, self.head_dim)
if freqs_cis is not None:
cos, sin = freqs_cis
query = apply_rotary_emb(query, cos, sin, is_neox_style=False)
key = apply_rotary_emb(key, cos, sin, is_neox_style=False)
# Self-attention only. FASTVIDEO_MLX_WINDOW=0/unset → full SDPA (byte-identical
# to the historical path). When >0, use chunked symmetric sliding-window
# attention (see windowed_attention.py). Cross-attn (attn2) stays full.
# Optional FASTVIDEO_MLX_WINDOW_SINK (default 0) adds global sink tokens.
q_bh = query.transpose(0, 2, 1, 3) # (B, H, S, D)
k_bh = key.transpose(0, 2, 1, 3)
v_bh = value.transpose(0, 2, 1, 3)
scale = self.head_dim**-0.5
window = int(os.environ.get("FASTVIDEO_MLX_WINDOW", "0") or "0")
if window > 0:
from fastvideo.mlx_runtime.windowed_attention import windowed_attention
_warn_windowed_attention_once(window)
sink = int(os.environ.get("FASTVIDEO_MLX_WINDOW_SINK", "0") or "0")
attn_output = windowed_attention(q_bh, k_bh, v_bh, window=window, sink=sink, scale=scale)
else:
attn_output = mx.fast.scaled_dot_product_attention(q_bh, k_bh, v_bh, scale=scale)
attn_output = attn_output.transpose(0, 2, 1, 3)
attn_output = attn_output.reshape(hidden_states.shape[0], -1, self.dim)
attn_output = linear(attn_output, self.weights["to_out.weight"], self.weights.get("to_out.bias"))
norm_hidden_states, hidden_states = scale_residual_layer_norm_scale_shift(
hidden_states,
attn_output,
gate_msa,
0.0,
0.0,
weight=self.weights["self_attn_residual_norm.norm.weight"],
bias=self.weights["self_attn_residual_norm.norm.bias"],
eps=self.eps,
)
norm_hidden_states = norm_hidden_states.astype(orig_dtype)
hidden_states = hidden_states.astype(orig_dtype)
attn_output = self.attn2(norm_hidden_states, encoder_hidden_states)
norm_hidden_states, hidden_states = scale_residual_layer_norm_scale_shift(
hidden_states,
attn_output,
1,
c_shift_msa,
c_scale_msa,
eps=self.eps,
)
norm_hidden_states = norm_hidden_states.astype(orig_dtype)
hidden_states = hidden_states.astype(orig_dtype)
ff_output = linear(norm_hidden_states, self.weights["ffn.fc_in.weight"], self.weights.get("ffn.fc_in.bias"))
ff_output = gelu_tanh(ff_output)
ff_output = linear(ff_output, self.weights["ffn.fc_out.weight"], self.weights.get("ffn.fc_out.bias"))
hidden_states = scale_residual(hidden_states, ff_output, c_gate_msa)
return hidden_states.astype(orig_dtype)
def mlx_block_weights_from_torch(torch_block) -> dict[str, mx.array]:
return {name: torch_to_mx(value) for name, value in torch_block.state_dict().items()}
class MLXWanDiT:
"""Experimental FP16 Wan/FastWan DiT forward path in MLX."""
def __init__(
self,
weights: dict[str, mx.array],
blocks: list[MLXWanTransformerBlock],
config: dict,
*,
compile: bool = False,
) -> None:
import os
self.weights = weights
self.blocks = blocks
self.config = config
self.num_heads = int(config["num_attention_heads"])
self.head_dim = int(config["attention_head_dim"])
self.hidden_size = self.num_heads * self.head_dim
self.ffn_dim = int(config["ffn_dim"])
self.in_channels = int(config["in_channels"])
self.out_channels = int(config["out_channels"])
self.patch_size = tuple(config["patch_size"])
self.freq_dim = int(config["freq_dim"])
# Opt-in graph fusion. With fixed weights and static shapes, the whole
# denoise-step forward is a pure function of (latents, timestep) -- a
# good mx.compile target. Off by default so the eager path stays the
# baseline; enable via constructor or FASTVIDEO_MLX_COMPILE=1 and verify
# with the benchmark's SSIM ~= 1.0 check.
self._enable_compile = compile or os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1"
self._compiled_forward: Callable[..., Any] | None = None
self._compiled_signature: tuple | None = None
def patch_embed(self, hidden_states):
batch, channels, frames, height, width = hidden_states.shape
pt, ph, pw = self.patch_size
patch_dim = channels * pt * ph * pw
x = hidden_states.reshape(batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw)
x = x.transpose(0, 2, 4, 6, 1, 3, 5, 7).reshape(batch, -1, patch_dim)
return linear(x, self.weights["patch_embedding.weight"], self.weights.get("patch_embedding.bias"))
def condition(self, timestep, encoder_hidden_states):
t_freq = timestep_embedding(timestep, self.freq_dim).astype(
weight_dtype(self.weights["condition_embedder.time_embedder.linear_1.weight"]))
temb = linear(
t_freq,
self.weights["condition_embedder.time_embedder.linear_1.weight"],
self.weights["condition_embedder.time_embedder.linear_1.bias"],
)
temb = silu(temb)
temb = linear(
temb,
self.weights["condition_embedder.time_embedder.linear_2.weight"],
self.weights["condition_embedder.time_embedder.linear_2.bias"],
)
timestep_proj = silu(temb)
timestep_proj = linear(
timestep_proj,
self.weights["condition_embedder.time_proj.weight"],
self.weights["condition_embedder.time_proj.bias"],
).reshape(timestep.shape[0], 6, self.hidden_size)
encoder_hidden_states = linear(
encoder_hidden_states,
self.weights["condition_embedder.text_embedder.linear_1.weight"],
self.weights["condition_embedder.text_embedder.linear_1.bias"],
)
encoder_hidden_states = gelu_tanh(encoder_hidden_states)
encoder_hidden_states = linear(
encoder_hidden_states,
self.weights["condition_embedder.text_embedder.linear_2.weight"],
self.weights["condition_embedder.text_embedder.linear_2.bias"],
)
return temb, timestep_proj, encoder_hidden_states
def output(self, hidden_states, temb, *, batch: int, frames: int, height: int, width: int):
pt, ph, pw = self.patch_size
post_patch_frames = frames // pt
post_patch_height = height // ph
post_patch_width = width // pw
shift, scale = mx_split_two(self.weights["scale_shift_table"] + temb[:, None, :], axis=1)
hidden_states = layer_norm(hidden_states, eps=float(self.config["eps"])) * (1.0 + scale) + shift
hidden_states = hidden_states.astype(weight_dtype(self.weights["proj_out.weight"]))
hidden_states = linear(hidden_states, self.weights["proj_out.weight"], self.weights["proj_out.bias"])
hidden_states = hidden_states.reshape(
batch,
post_patch_frames,
post_patch_height,
post_patch_width,
pt,
ph,
pw,
self.out_channels,
)
hidden_states = hidden_states.transpose(0, 7, 1, 4, 2, 5, 3, 6)
return hidden_states.reshape(batch, self.out_channels, frames, height, width)
def _forward(self, hidden_states, encoder_hidden_states, timestep, cos, sin):
"""Pure forward used both eagerly and as the mx.compile target.
``cos``/``sin`` are passed as separate array args (rather than a tuple)
so the function traces cleanly under mx.compile.
"""
batch, _, frames, height, width = hidden_states.shape
freqs_cis = (cos, sin) if cos is not None else None
hidden_states = self.patch_embed(hidden_states)
temb, timestep_proj, encoder_hidden_states = self.condition(timestep, encoder_hidden_states)
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, freqs_cis=freqs_cis)
return self.output(hidden_states, temb, batch=batch, frames=frames, height=height, width=width)
def __call__(self, hidden_states, encoder_hidden_states, timestep, freqs_cis):
cos, sin = freqs_cis if freqs_cis is not None else (None, None)
if self._enable_compile and cos is not None:
import mlx.core as mx
# mx.compile keeps one traced graph per input signature, and each
# graph pins its own materialization of the quantized weights. The
# two-pass modes (--refine) call the DiT at a second resolution, so
# keeping both graphs alive doubles resident DiT memory: 14B refine
# peaked at 34.7 GiB instead of 20.8 GiB. Retire the previous graph
# when the signature changes; the retrace costs far less than a
# second copy of the weights.
signature = (hidden_states.shape, encoder_hidden_states.shape, timestep.shape)
if self._compiled_forward is not None and signature != self._compiled_signature:
self._compiled_forward = None
self._compiled_signature = None
mx.clear_cache()
if self._compiled_forward is None:
self._compiled_forward = mx.compile(self._forward)
self._compiled_signature = signature
compiled_forward = self._compiled_forward
try:
return compiled_forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
except Exception as exc: # noqa: BLE001 - some quant graphs may not trace; fall back to eager.
logger.warning("mx.compile forward failed (%s); falling back to eager execution.", exc)
self._enable_compile = False
self._compiled_forward = None
self._compiled_signature = None
return self._forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
def mx_split_two(x, *, axis: int):
import mlx.core as mx
left, right = mx.split(x, 2, axis=axis)
return left, right
def _load_safetensor_value(handle, name: str):
return handle.get_tensor(name)
def _load_mx_array_from_safetensor(handle, name: str, dtype):
"""Load a safetensors value and cast before creating the MLX array.
The FastWan Diffusers checkpoint is fp32. Creating an MLX array first and
then casting it to fp16 briefly materializes a large fp32 MLX allocation.
Casting the CPU tensor before crossing into MLX keeps the transient GPU-side
footprint lower.
"""
import mlx.core as mx
import torch
tensor = handle.get_tensor(name)
if dtype == mx.float16:
tensor = tensor.to(torch.float16)
elif dtype == mx.float32:
tensor = tensor.to(torch.float32)
elif dtype == mx.bfloat16:
# NumPy has no bfloat16, so bridge through fp32 and cast on-device below.
tensor = tensor.to(torch.float32)
array = mx.array(tensor.numpy())
del tensor
if dtype is not None and array.dtype != dtype:
array = array.astype(dtype)
mx.eval(array)
return array
def _eval_loaded_weight(value) -> None:
import mlx.core as mx
if isinstance(value, QuantizedMatrix):
eval_args = [value.weight, value.scales]
if value.biases is not None:
eval_args.append(value.biases)
mx.eval(*eval_args)
else:
mx.eval(value)
# Diffusers-to-FastVideo key mapping for WanTransformerBlock weights.
# Shared by both MLX and torch block loaders to keep mappings synchronized.
_WAN_BLOCK_KEY_MAP = {
"scale_shift_table": "scale_shift_table",
"attn1.to_q.weight": "to_q.weight",
"attn1.to_q.bias": "to_q.bias",
"attn1.to_k.weight": "to_k.weight",
"attn1.to_k.bias": "to_k.bias",
"attn1.to_v.weight": "to_v.weight",
"attn1.to_v.bias": "to_v.bias",
"attn1.to_out.0.weight": "to_out.weight",
"attn1.to_out.0.bias": "to_out.bias",
"attn1.norm_q.weight": "norm_q.weight",
"attn1.norm_k.weight": "norm_k.weight",
"attn2.to_q.weight": "attn2.to_q.weight",
"attn2.to_q.bias": "attn2.to_q.bias",
"attn2.to_k.weight": "attn2.to_k.weight",
"attn2.to_k.bias": "attn2.to_k.bias",
"attn2.to_v.weight": "attn2.to_v.weight",
"attn2.to_v.bias": "attn2.to_v.bias",
"attn2.to_out.0.weight": "attn2.to_out.weight",
"attn2.to_out.0.bias": "attn2.to_out.bias",
"attn2.norm_q.weight": "attn2.norm_q.weight",
"attn2.norm_k.weight": "attn2.norm_k.weight",
"ffn.net.0.proj.weight": "ffn.fc_in.weight",
"ffn.net.0.proj.bias": "ffn.fc_in.bias",
"ffn.net.2.weight": "ffn.fc_out.weight",
"ffn.net.2.bias": "ffn.fc_out.bias",
"norm2.weight": "self_attn_residual_norm.norm.weight",
"norm2.bias": "self_attn_residual_norm.norm.bias",
}
def mlx_block_weights_from_diffusers_safetensors(
checkpoint_path: str | Path,
*,
block_index: int = 0,
quantization: str | MLXQuantizationSpec | None = None,
dtype=None,
) -> dict[str, mx.array]:
"""Load one Diffusers-format Wan block into the MLX dense-block key layout."""
from safetensors import safe_open
prefix = f"blocks.{block_index}."
key_map = _WAN_BLOCK_KEY_MAP
spec = MLXQuantizationSpec.from_name(quantization) if (quantization is None
or isinstance(quantization, str)) else quantization
ensure_quantization_supported(spec)
matrix_targets = {target for target in key_map.values() if target.endswith(".weight") and "norm" not in target}
weights = {}
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for source_name, target_name in key_map.items():
full = prefix + source_name
if full not in available:
# Biases are optional: e.g. Wan2.1-14B has bias-free attention/FFN.
# The block forward already fetches biases via ``.get(...)``.
if source_name.endswith(".bias"):
continue
raise KeyError(f"missing required block weight: {full}")
array = _load_mx_array_from_safetensor(handle, full, dtype)
loaded = quantize_matrix(array, spec) if target_name in matrix_targets else array
_eval_loaded_weight(loaded)
weights[target_name] = loaded
del array
return weights
def mlx_dit_from_diffusers_safetensors(
checkpoint_path: str | Path,
config_path: str | Path,
*,
dtype: str = "fp16",
num_blocks: int | None = None,
quantization: str | MLXQuantizationSpec | None = None,
compile: bool = False,
) -> MLXWanDiT:
import mlx.core as mx
from safetensors import safe_open
config = json.loads(Path(config_path).read_text())
total_blocks = int(config["num_layers"])
if num_blocks is None:
num_blocks = total_blocks
cast_dtype = {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[dtype]
spec = MLXQuantizationSpec.from_name(quantization) if (quantization is None
or isinstance(quantization, str)) else quantization
ensure_quantization_supported(spec)
top_level_names = [
"patch_embedding.weight",
"patch_embedding.bias",
"condition_embedder.time_embedder.linear_1.weight",
"condition_embedder.time_embedder.linear_1.bias",
"condition_embedder.time_embedder.linear_2.weight",
"condition_embedder.time_embedder.linear_2.bias",
"condition_embedder.time_proj.weight",
"condition_embedder.time_proj.bias",
"condition_embedder.text_embedder.linear_1.weight",
"condition_embedder.text_embedder.linear_1.bias",
"condition_embedder.text_embedder.linear_2.weight",
"condition_embedder.text_embedder.linear_2.bias",
"scale_shift_table",
"proj_out.weight",
"proj_out.bias",
]
weights = {}
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for name in top_level_names:
if name not in available:
if name.endswith(".bias"):
continue
raise KeyError(f"missing required weight: {name}")
array = _load_mx_array_from_safetensor(handle, name, cast_dtype)
if name == "patch_embedding.weight":
array = array.reshape(int(config["num_attention_heads"]) * int(config["attention_head_dim"]), -1)
if name.endswith(".weight") and name not in {"scale_shift_table"}:
loaded = quantize_matrix(array, spec)
else:
loaded = array
_eval_loaded_weight(loaded)
weights[name] = loaded
del array
blocks = []
for block_index in range(num_blocks):
block_weights = mlx_block_weights_from_diffusers_safetensors(
checkpoint_path,
block_index=block_index,
quantization=spec,
dtype=cast_dtype,
)
block_weights = {
name: (value if isinstance(value, QuantizedMatrix) else value.astype(cast_dtype))
for name, value in block_weights.items()
}
for value in block_weights.values():
_eval_loaded_weight(value)
blocks.append(
MLXWanTransformerBlock(
block_weights,
dim=int(config["num_attention_heads"]) * int(config["attention_head_dim"]),
ffn_dim=int(config["ffn_dim"]),
num_heads=int(config["num_attention_heads"]),
eps=float(config["eps"]),
))
return MLXWanDiT(weights, blocks, config, compile=compile)
def torch_block_state_from_diffusers_safetensors(
checkpoint_path: str | Path,
*,
block_index: int = 0,
) -> dict[str, torch.Tensor]:
"""Load one Diffusers-format Wan block into FastVideo's dense block keys."""
from safetensors import safe_open
prefix = f"blocks.{block_index}."
key_map = _WAN_BLOCK_KEY_MAP
state = {}
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for source_name, target_name in key_map.items():
full = prefix + source_name
if full not in available:
# Biases are optional: e.g. Wan2.1-14B has bias-free attention/FFN.
if source_name.endswith(".bias"):
continue
raise KeyError(f"missing required block weight: {full}")
state[target_name] = handle.get_tensor(full).float()
return state
+154
View File
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
"""Pixel-space spatial resampling for decoded MLX Wan frames.
Spatial fast mode denoises on a smaller latent grid and has to get back to
the requested output size. That resize belongs *here* — after the VAE
decode — and not in latent space.
A Wan latent cell is a learned code for an 8x8 (Wan2.1) or 16x16 (Wan2.2)
pixel block, not a low-pass sample of the image. Linearly blending two
adjacent codes does not produce the code of the blended blocks; it produces
a vector the decoder was never trained on. The decoder answers with smeared,
ringing texture laid over otherwise-correct structure — the silhouette
survives, the detail turns to haze. Measured on Wan2.1-1.3B at 480x832, a
2x bilinear latent upsample destroys 62% of the latent's high-frequency
energy while leaving its overall magnitude intact, which is exactly the
signature of that veil.
Resampling decoded RGB frames has none of that problem: an image *is* a
sampled 2-D signal, so Lanczos/cubic interpolation is the operation it was
defined for. The result is soft — it carries stage-1's real detail budget
and no more — but it is clean and coherent.
"""
from __future__ import annotations
from collections.abc import Iterable
import numpy as np
# Pixel-space interpolation kernels, best-quality first. ``lanczos`` is the
# default: it holds edges better than cubic at 2x with no visible ringing on
# decoder output, which is already band-limited.
PIXEL_UPSAMPLE_MODES = ("lanczos", "cubic", "bilinear", "nearest")
DEFAULT_PIXEL_UPSAMPLE_MODE = "lanczos"
def _interpolation_flag(mode: str) -> int:
"""
Map a pixel upsample mode name onto its OpenCV interpolation flag.
Parameters:
mode (str): One of :data:`PIXEL_UPSAMPLE_MODES`.
Returns:
int: The matching ``cv2.INTER_*`` flag.
Raises:
ValueError: If the mode is not a supported pixel upsample mode.
"""
import cv2
flags = {
"lanczos": cv2.INTER_LANCZOS4,
"cubic": cv2.INTER_CUBIC,
"bilinear": cv2.INTER_LINEAR,
"nearest": cv2.INTER_NEAREST,
}
try:
return flags[mode]
except KeyError:
raise ValueError(f"Unsupported pixel upsample mode: {mode!r} "
f"(expected one of {', '.join(PIXEL_UPSAMPLE_MODES)})") from None
def unsharp(frame: np.ndarray, amount: float) -> np.ndarray:
"""Light unsharp mask, used to counter resampling / optical-flow softening.
Parameters:
frame (np.ndarray): HxWx3 uint8 RGB frame.
amount (float): Strength; ``0`` returns the frame unchanged.
Returns:
np.ndarray: A new frame; the input is never modified in place.
"""
if amount <= 0.0:
return frame
import cv2
blur = cv2.GaussianBlur(frame, (0, 0), 1.0)
return cv2.addWeighted(frame, 1.0 + amount, blur, -amount, 0)
def upsample_frame(
frame: np.ndarray,
*,
width: int,
height: int,
mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
sharpen: float = 0.0,
) -> np.ndarray:
"""
Resample one decoded RGB frame to the target pixel size.
Parameters:
frame (np.ndarray): HxWx3 uint8 RGB frame.
width (int): Target width in pixels.
height (int): Target height in pixels.
mode (str): Interpolation kernel, one of :data:`PIXEL_UPSAMPLE_MODES`.
sharpen (float): Unsharp strength applied after the resize.
Returns:
np.ndarray: A new frame at ``height x width``; already-correct sizes
are still passed through ``sharpen``.
Raises:
ValueError: If the frame is not HxWx3, or the target size is not positive.
"""
import cv2
array = np.asarray(frame)
if array.ndim != 3 or array.shape[2] != 3:
raise ValueError(f"frame must have shape HxWx3, got {array.shape}")
if width <= 0 or height <= 0:
raise ValueError(f"target size must be positive, got {width}x{height}")
if array.dtype != np.uint8:
array = np.clip(array, 0, 255).astype(np.uint8)
if (array.shape[0], array.shape[1]) != (height, width):
array = cv2.resize(array, (width, height), interpolation=_interpolation_flag(mode))
return unsharp(array, sharpen)
def upsample_frames(
frames: Iterable[np.ndarray],
*,
width: int,
height: int,
mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
sharpen: float = 0.0,
) -> list[np.ndarray]:
"""
Resample every decoded frame to the target pixel size.
Parameters:
frames (Iterable[np.ndarray]): Decoded HxWx3 uint8 RGB frames.
width (int): Target width in pixels.
height (int): Target height in pixels.
mode (str): Interpolation kernel, one of :data:`PIXEL_UPSAMPLE_MODES`.
sharpen (float): Unsharp strength applied after each resize.
Returns:
list[np.ndarray]: New frames at the target size, in input order.
"""
return [upsample_frame(frame, width=width, height=height, mode=mode, sharpen=sharpen) for frame in frames]
__all__ = [
"DEFAULT_PIXEL_UPSAMPLE_MODE",
"PIXEL_UPSAMPLE_MODES",
"unsharp",
"upsample_frame",
"upsample_frames",
]
+243
View File
@@ -0,0 +1,243 @@
# SPDX-License-Identifier: Apache-2.0
"""Memory-tier helpers for Apple Silicon MLX/MPS experiments.
macOS does not expose a perfect "pretend this machine only has 16 GB unified
memory" switch. MLX can cap the allocator used by the Apple-native DiT path,
and PyTorch MPS exposes process-level watermark environment variables for the
hybrid prompt/decode stages. Applying both gives benchmark and generation
entrypoints a practical, explicit way to exercise memory-tier presets.
"""
from __future__ import annotations
import argparse
import gc
import os
from dataclasses import dataclass, field
from typing import Any
GIB = 1024**3
@dataclass(frozen=True)
class AppliedMemoryLimits:
"""Memory limits applied for one Apple Silicon benchmark/generation process."""
mlx_memory_limit_gib: float | None = None
mlx_cache_limit_gib: float | None = None
mlx_disable_cache: bool = False
mlx_wired_limit_gib: float | None = None
torch_mps_high_watermark_ratio: float | None = None
torch_mps_low_watermark_ratio: float | None = None
applied_bytes: dict[str, int] = field(default_factory=dict)
previous_bytes: dict[str, int] = field(default_factory=dict)
errors: dict[str, str] = field(default_factory=dict)
def as_metrics(self) -> dict[str, int | float | str | bool | None]:
"""Flatten the configured memory limits, applied values, previous values, and errors into a metrics dictionary.
Returns:
dict[str, int | float | str | bool | None]: Metrics keyed by limit names and their corresponding values.
"""
metrics: dict[str, int | float | str | bool | None] = {
"mlx_memory_limit_gib": self.mlx_memory_limit_gib,
"mlx_cache_limit_gib": self.mlx_cache_limit_gib,
"mlx_disable_cache": self.mlx_disable_cache,
"mlx_wired_limit_gib": self.mlx_wired_limit_gib,
"torch_mps_high_watermark_ratio": self.torch_mps_high_watermark_ratio,
"torch_mps_low_watermark_ratio": self.torch_mps_low_watermark_ratio,
}
for name, value in self.applied_bytes.items():
metrics[f"{name}_bytes"] = value
for name, value in self.previous_bytes.items():
metrics[f"previous_{name}_bytes"] = value
for name, error in self.errors.items():
metrics[f"{name}_error"] = error
return metrics
def gib_to_bytes(value: float | None) -> int | None:
"""
Convert a positive memory limit from GiB to bytes.
Parameters:
value (float | None): Memory limit in GiB, or `None` when unset.
Returns:
int | None: The memory limit in bytes, or `None` when no limit is provided.
Raises:
ValueError: If `value` is zero or negative.
"""
if value is None:
return None
if value <= 0:
raise ValueError(f"Memory limit must be positive GiB, got {value}")
return int(value * GIB)
def cleanup_mlx(mx_module: Any | None = None) -> None:
"""Collect unreachable MLX objects, then release their allocator cache."""
if mx_module is None:
import mlx.core as mx
mx_module = mx
gc.collect()
mx_module.clear_cache()
def cleanup_torch_mps(torch_module: Any | None = None) -> None:
"""Collect unreachable Torch objects, then release the MPS allocator cache."""
if torch_module is None:
import torch
torch_module = torch
gc.collect()
if torch_module.backends.mps.is_available():
torch_module.mps.empty_cache()
def _set_mps_env(name: str, value: float | None) -> float | None:
"""Set a PyTorch MPS watermark environment variable.
Parameters:
name (str): Name of the environment variable to set.
value (float | None): Watermark ratio, or `None` to leave the variable unchanged.
Returns:
float | None: The configured watermark ratio, or `None` when no value is provided.
Raises:
ValueError: If `value` is negative.
"""
if value is None:
return None
if value < 0:
raise ValueError(f"{name} must be non-negative, got {value}")
os.environ[name] = str(value)
return value
def apply_memory_limits(
*,
mlx_memory_limit_gib: float | None = None,
mlx_cache_limit_gib: float | None = None,
mlx_disable_cache: bool = False,
mlx_wired_limit_gib: float | None = None,
torch_mps_high_watermark_ratio: float | None = None,
torch_mps_low_watermark_ratio: float | None = None,
mx_module: Any | None = None,
) -> AppliedMemoryLimits:
"""Apply optional MLX allocator limits and PyTorch MPS watermarks.
PyTorch reads MPS watermark variables when the MPS backend initializes, so
call this before importing PyTorch. Specifying only a high watermark sets the
low watermark to ``0.0``. MLX limit-setting failures are recorded in the
result and do not prevent other limits from being applied.
Parameters:
mlx_memory_limit_gib (float | None): Maximum MLX memory in GiB.
mlx_cache_limit_gib (float | None): Maximum MLX cache size in GiB.
mlx_disable_cache (bool): Whether to disable the MLX cache.
mlx_wired_limit_gib (float | None): Maximum MLX wired memory in GiB.
torch_mps_high_watermark_ratio (float | None): PyTorch MPS high watermark
ratio.
torch_mps_low_watermark_ratio (float | None): PyTorch MPS low watermark
ratio.
Returns:
AppliedMemoryLimits: Configured values, applied and previous MLX byte
limits, MPS watermark values, and per-limit errors.
"""
if torch_mps_high_watermark_ratio is not None and torch_mps_low_watermark_ratio is None:
torch_mps_low_watermark_ratio = 0.0
high = _set_mps_env("PYTORCH_MPS_HIGH_WATERMARK_RATIO", torch_mps_high_watermark_ratio)
low = _set_mps_env("PYTORCH_MPS_LOW_WATERMARK_RATIO", torch_mps_low_watermark_ratio)
memory_bytes = gib_to_bytes(mlx_memory_limit_gib)
cache_bytes = 0 if mlx_disable_cache else gib_to_bytes(mlx_cache_limit_gib)
wired_bytes = gib_to_bytes(mlx_wired_limit_gib)
applied: dict[str, int] = {}
previous: dict[str, int] = {}
errors: dict[str, str] = {}
if memory_bytes is not None or cache_bytes is not None or wired_bytes is not None:
if mx_module is None:
import mlx.core as mx
mx_module = mx
# Apply each limit independently; record failures without stopping.
limits = [
("mlx_memory_limit", memory_bytes, mx_module.set_memory_limit),
("mlx_cache_limit", cache_bytes, mx_module.set_cache_limit),
("mlx_wired_limit", wired_bytes, mx_module.set_wired_limit),
]
for name, value, setter in limits:
if value is not None:
try:
previous[name] = int(setter(value))
applied[name] = value
except Exception as exc: # noqa: BLE001 - macOS/system-limit dependent.
errors[name] = f"{type(exc).__name__}: {exc}"
return AppliedMemoryLimits(
mlx_memory_limit_gib=mlx_memory_limit_gib,
mlx_cache_limit_gib=mlx_cache_limit_gib,
mlx_disable_cache=mlx_disable_cache,
mlx_wired_limit_gib=mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=high,
torch_mps_low_watermark_ratio=low,
applied_bytes=applied,
previous_bytes=previous,
errors=errors,
)
def add_memory_limit_args(
parser: argparse.ArgumentParser,
*,
mlx_memory_limit_gib: float | None = None,
mlx_cache_limit_gib: float | None = None,
mlx_disable_cache: bool = False,
mlx_wired_limit_gib: float | None = None,
torch_mps_high_watermark_ratio: float | None = None,
torch_mps_low_watermark_ratio: float | None = None,
) -> None:
"""
Add configurable Apple Silicon memory-limit options to an argument parser.
Parameters:
parser (argparse.ArgumentParser): Parser to which the options are added.
mlx_memory_limit_gib (float | None): Default MLX memory limit in GiB.
mlx_cache_limit_gib (float | None): Default MLX cache limit in GiB.
mlx_disable_cache (bool): Whether the cache limit defaults to zero.
mlx_wired_limit_gib (float | None): Default MLX wired-memory limit in GiB.
torch_mps_high_watermark_ratio (float | None): Default PyTorch MPS high-watermark ratio.
torch_mps_low_watermark_ratio (float | None): Default PyTorch MPS low-watermark ratio.
"""
parser.add_argument("--mlx-memory-limit-gib",
type=float,
default=mlx_memory_limit_gib,
help="Set MLX memory limit in GiB for memory-tier testing (DiT path).")
parser.add_argument("--mlx-cache-limit-gib",
type=float,
default=mlx_cache_limit_gib,
help="Set MLX cache limit in GiB. Use --mlx-disable-cache to force 0.")
parser.add_argument("--mlx-disable-cache",
action="store_true",
default=mlx_disable_cache,
help="Set MLX cache limit to 0 for stricter memory-tier tests.")
parser.add_argument("--mlx-wired-limit-gib",
type=float,
default=mlx_wired_limit_gib,
help="Set MLX wired-memory limit in GiB where supported by macOS/MLX.")
parser.add_argument("--torch-mps-high-watermark-ratio",
type=float,
default=torch_mps_high_watermark_ratio,
help="Set PYTORCH_MPS_HIGH_WATERMARK_RATIO before importing torch.")
parser.add_argument("--torch-mps-low-watermark-ratio",
type=float,
default=torch_mps_low_watermark_ratio,
help="Set PYTORCH_MPS_LOW_WATERMARK_RATIO before importing torch.")
+129
View File
@@ -0,0 +1,129 @@
# SPDX-License-Identifier: Apache-2.0
"""Best-effort prompt-embedding cache shared by the MLX entrypoints."""
from __future__ import annotations
import hashlib
import io
import json
import logging
import tempfile
from pathlib import Path
import numpy as np
logger = logging.getLogger(__name__)
def fingerprint_digest(fingerprint: dict[str, object]) -> str:
payload = json.dumps(fingerprint, sort_keys=True, separators=(",", ":"))
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def text_encoder_fingerprint(model_root: Path) -> dict[str, object]:
"""Return a cheap identity for the tokenizer and text-encoder files."""
root = model_root.resolve()
components = [path for name in ("tokenizer", "text_encoder") if (path := root / name).is_dir()]
scan_roots = components or [root]
files: list[list[object]] = []
complete = True
try:
for scan_root in scan_roots:
for path in sorted(scan_root.rglob("*")):
try:
if not path.is_file():
continue
stat = path.stat()
files.append([
path.relative_to(root).as_posix(),
stat.st_size,
stat.st_mtime_ns,
stat.st_ctime_ns,
])
except OSError:
complete = False
except OSError:
complete = False
# ponytail: metadata avoids hashing multi-GB weights; use a model manifest
# if supported workflows ever preserve size, mtime, and ctime while mutating.
return {"root": str(root), "files": files, "complete": complete}
def prompt_cache_meta_path(cache_path: Path) -> Path:
return cache_path.with_suffix(cache_path.suffix + ".json")
def _fingerprint_is_complete(fingerprint: dict[str, object]) -> bool:
text_encoder = fingerprint.get("text_encoder")
return not isinstance(text_encoder, dict) or text_encoder.get("complete") is not False
def load_prompt_cache(
cache_path: Path | None,
fingerprint: dict[str, object],
) -> np.ndarray | None:
"""Load a matching cache entry, treating every cache failure as a miss."""
if cache_path is None or not _fingerprint_is_complete(fingerprint):
return None
try:
metadata = json.loads(prompt_cache_meta_path(cache_path).read_text())
if not isinstance(metadata, dict):
return None
if metadata.get("fingerprint_sha256") != fingerprint_digest(fingerprint):
return None
payload = cache_path.read_bytes()
if metadata.get("data_sha256") != hashlib.sha256(payload).hexdigest():
return None
array = np.load(io.BytesIO(payload), allow_pickle=False)
return array if isinstance(array, np.ndarray) else None
except (EOFError, OSError, UnicodeError, ValueError) as exc:
logger.info("Prompt cache read skipped for %s: %s", cache_path, exc)
return None
def _atomic_write(path: Path, payload: bytes) -> None:
temp_path: Path | None = None
try:
with tempfile.NamedTemporaryFile(
mode="wb",
dir=path.parent,
prefix=f".{path.name}.",
suffix=".tmp",
delete=False,
) as handle:
temp_path = Path(handle.name)
handle.write(payload)
temp_path.replace(path)
finally:
if temp_path is not None:
temp_path.unlink(missing_ok=True)
def save_prompt_cache(
cache_path: Path | None,
embeds: np.ndarray,
fingerprint: dict[str, object],
) -> bool:
"""Atomically publish an integrity-bound cache entry when possible."""
if cache_path is None or not _fingerprint_is_complete(fingerprint):
return False
try:
cache_path.parent.mkdir(parents=True, exist_ok=True)
buffer = io.BytesIO()
np.save(buffer, np.asarray(embeds), allow_pickle=False)
payload = buffer.getvalue()
metadata = (json.dumps(
{
"fingerprint_sha256": fingerprint_digest(fingerprint),
"data_sha256": hashlib.sha256(payload).hexdigest(),
"fingerprint": fingerprint,
},
indent=2) + "\n").encode("utf-8")
# Publish data first. Until metadata follows, old metadata's data digest
# makes the torn pair a harmless miss rather than a stale cache hit.
_atomic_write(cache_path, payload)
_atomic_write(prompt_cache_meta_path(cache_path), metadata)
return True
except (OSError, TypeError, ValueError) as exc:
logger.info("Prompt cache write skipped for %s: %s", cache_path, exc)
return False
+454
View File
@@ -0,0 +1,454 @@
# SPDX-License-Identifier: Apache-2.0
"""Local prompt enrichment for the MLX Wan runtime (H3 Context-IR-style).
Wan's training captions are long and cinematic; short user prompts leave
quality on the table. This module expands a raw prompt into Wan-style
shot language **on device** — no remote API, no training.
Backends (first match wins):
1. **mlx-lm** — optional local LLM (``--enhance-prompt-model``).
2. **template** — deterministic cinematic expansion (always available).
System-prompt contract matches the streaming server's enhancer defaults
in ``fastvideo/entrypoints/streaming/prompt/enhancer.py`` so remote and
local paths stay interchangeable.
"""
from __future__ import annotations
import hashlib
import json
import re
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Keep in lockstep with streaming PromptEnhancer defaults (enhance op).
DEFAULT_ENHANCE_SYSTEM_PROMPT = ("You are a prompt enhancer for cinematic video generation. Given "
"a user prompt, produce an enhanced prompt that is more vivid, "
"specific, and concrete. Keep the subject intact; add lighting, "
"camera, and motion detail. Reply with just the enhanced prompt.")
# Small default that fits 16 GB Macs alongside the 1.3B DiT when the user
# opts into mlx-lm. Override with --enhance-prompt-model.
DEFAULT_MLX_LM_MODEL = "mlx-community/Qwen2.5-0.5B-Instruct-4bit"
_CAMERA_CUES = (
"cinematic",
"camera",
"lens",
"shot",
"bokeh",
"dolly",
"tracking",
"close-up",
"wide shot",
"handheld",
"steadicam",
)
_LIGHT_CUES = (
"light",
"lighting",
"sun",
"golden hour",
"neon",
"rim light",
"softbox",
"overcast",
"moonlight",
"volumetric",
)
_MOTION_CUES = (
"moving",
"motion",
"walk",
"run",
"flies",
"flying",
"drifts",
"sails",
"flows",
"pan",
"tilt",
"zoom",
)
@dataclass(frozen=True)
class EnhanceResult:
"""Outcome of a prompt enrichment call."""
original: str
enhanced: str
backend: str
elapsed_s: float
model: str | None = None
@property
def changed(self) -> bool:
"""Indicates whether the enhanced prompt differs from the original after trimming surrounding whitespace.
Returns:
bool: `True` if the prompts differ, `False` otherwise.
"""
return self.enhanced.strip() != self.original.strip()
def _normalize_user_prompt(prompt: str) -> str:
"""
Normalize a user prompt for enhancement.
Parameters:
prompt (str): User-provided prompt text.
Returns:
str: The prompt with leading and trailing whitespace removed and internal whitespace collapsed.
Raises:
ValueError: If the prompt is empty after whitespace normalization.
"""
text = " ".join(prompt.strip().split())
if not text:
raise ValueError("prompt must be non-empty")
return text
def _already_rich(prompt: str) -> bool:
"""
Determine whether a prompt already contains substantial camera and lighting detail.
Returns:
bool: `true` if the prompt is at least 160 characters long and includes camera and lighting cues, `false` otherwise.
"""
lower = prompt.lower()
has_camera = any(c in lower for c in _CAMERA_CUES)
has_light = any(c in lower for c in _LIGHT_CUES)
return len(prompt) >= 160 and has_camera and has_light
def enhance_prompt_template(prompt: str) -> str:
"""
Expand a prompt with cinematic camera, lighting, motion, and visual-quality details.
Rich prompts are preserved, while thinner prompts receive deterministic enhancements
without changing their subject.
Returns:
str: The original or expanded prompt with normalized whitespace and punctuation.
"""
text = _normalize_user_prompt(prompt)
if _already_rich(text):
return text
lower = text.lower()
parts = [text.rstrip(".")]
if not any(c in lower for c in _CAMERA_CUES):
parts.append("shot on a 35mm anamorphic lens, gentle handheld micro-movement, "
"shallow depth of field")
if not any(c in lower for c in _LIGHT_CUES):
parts.append("natural cinematic lighting with soft volumetric haze and subtle "
"rim light separating subject from background")
if not any(c in lower for c in _MOTION_CUES):
parts.append("smooth continuous motion with grounded physics")
parts.append("highly detailed, coherent temporal continuity, film grain, "
"color graded like a contemporary drama")
enhanced = ", ".join(parts)
# Single trailing period; collapse duplicate whitespace.
enhanced = re.sub(r"\s+", " ", enhanced).strip()
if not enhanced.endswith("."):
enhanced += "."
return enhanced
def enhance_prompt_mlx_lm(
prompt: str,
*,
model: str = DEFAULT_MLX_LM_MODEL,
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
max_tokens: int = 128,
temp: float = 0.6,
) -> str:
"""
Enhance a user prompt with a locally hosted mlx-lm instruction model.
Parameters:
prompt (str): The prompt to enhance.
model (str): The mlx-lm model identifier or path.
system_prompt (str): Instructions that guide prompt enhancement.
max_tokens (int): Maximum number of tokens to generate.
temp (float): Sampling temperature for generation.
Returns:
str: The enhanced prompt.
Raises:
RuntimeError: If mlx-lm is unavailable or produces an empty result.
"""
try:
from mlx_lm import generate, load
except ImportError as exc: # pragma: no cover - optional dep
raise RuntimeError("mlx-lm is not installed. `uv pip install mlx-lm` or use "
"--enhance-prompt-backend template.") from exc
text = _normalize_user_prompt(prompt)
logger.info("[MLX enhance] loading %s", model)
mlx_model, tokenizer = load(model)
messages = [
{
"role": "system",
"content": system_prompt
},
{
"role": "user",
"content": text
},
]
if hasattr(tokenizer, "apply_chat_template"):
chat = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
else: # pragma: no cover - ancient tokenizers
chat = f"{system_prompt}\n\nUser: {text}\nAssistant:"
raw = generate(
mlx_model,
tokenizer,
prompt=chat,
max_tokens=max_tokens,
temp=temp,
verbose=False,
)
enhanced = _clean_llm_output(raw, original=text)
if not enhanced:
raise RuntimeError("mlx-lm returned an empty enhance result")
return enhanced
def _clean_llm_output(raw: str, *, original: str) -> str:
"""
Clean generated prompt text and fall back to the original when the result is too short.
Parameters:
raw (str): Raw text produced by the language model.
original (str): Original prompt used as the fallback value.
Returns:
str: Cleaned first paragraph of the generated text, or the original prompt when the generated text is too short.
"""
text = raw.strip()
# Drop common prefatory phrases.
for prefix in (
"enhanced prompt:",
"here's the enhanced prompt:",
"here is the enhanced prompt:",
"sure:",
"sure,",
):
if text.lower().startswith(prefix):
text = text[len(prefix):].strip()
# Keep first non-empty paragraph only.
para = text.split("\n\n")[0].strip()
para = " ".join(para.split())
if len(para) < max(12, len(original) // 4):
return original
return para
def enhance_prompt(
prompt: str,
*,
backend: str = "auto",
model: str | None = None,
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
max_tokens: int = 128,
) -> EnhanceResult:
"""Enhance a prompt using the selected backend, falling back to a deterministic template when configured for automatic selection.
Parameters:
prompt (str): The prompt to enhance.
backend (str): The enhancement backend: ``"auto"``, ``"mlx-lm"``, or ``"template"``.
model (str | None): The MLX language model to use.
system_prompt (str): Instructions provided to the MLX language model.
max_tokens (int): Maximum number of tokens generated by the MLX language model.
Returns:
EnhanceResult: The original and enhanced prompts, selected backend, timing information, and model metadata.
Raises:
ValueError: If the prompt is empty or the backend is unsupported.
Exception: If the explicitly selected ``"mlx-lm"`` backend fails.
"""
text = _normalize_user_prompt(prompt)
backend_norm = (backend or "auto").lower()
if backend_norm not in {"auto", "mlx-lm", "template"}:
raise ValueError(f"Unknown enhance backend: {backend}")
start = time.perf_counter()
used_model: str | None = None
if backend_norm in {"auto", "mlx-lm"}:
try:
used_model = model or DEFAULT_MLX_LM_MODEL
enhanced = enhance_prompt_mlx_lm(
text,
model=used_model,
system_prompt=system_prompt,
max_tokens=max_tokens,
)
return EnhanceResult(
original=text,
enhanced=enhanced,
backend="mlx-lm",
elapsed_s=time.perf_counter() - start,
model=used_model,
)
except Exception as exc:
if backend_norm == "mlx-lm":
raise
logger.info(
"[MLX enhance] mlx-lm unavailable (%s); using template backend",
exc,
)
enhanced = enhance_prompt_template(text)
return EnhanceResult(
original=text,
enhanced=enhanced,
backend="template",
elapsed_s=time.perf_counter() - start,
model=None,
)
def enhance_cache_path(
prompt: str,
*,
backend: str,
model: str | None,
cache_dir: Path | None = None,
) -> Path:
"""Content-addressed cache file for an enhanced prompt string."""
root = cache_dir or (Path.home() / ".cache" / "fastvideo" / "enhanced_prompts")
key = "\0".join([prompt, backend, model or ""])
digest = hashlib.sha256(key.encode("utf-8")).hexdigest()[:24]
return root / f"{digest}.json"
def load_or_enhance_prompt(
prompt: str,
*,
backend: str = "auto",
model: str | None = None,
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
max_tokens: int = 128,
cache: bool = True,
cache_dir: Path | None = None,
) -> EnhanceResult:
"""
Enhance a prompt, reusing a cached result when available.
Parameters:
prompt (str): The prompt to enhance.
backend (str): Enhancement backend to use.
model (str | None): Optional model identifier.
system_prompt (str): System prompt for model-based enhancement.
max_tokens (int): Maximum number of tokens generated by the model.
cache (bool): Whether to read and write the on-disk cache.
cache_dir (Path | None): Optional directory for cached results.
Returns:
EnhanceResult: The enhanced prompt and backend metadata. Cached results are marked with the ``"cache"`` backend.
"""
text = _normalize_user_prompt(prompt)
path = enhance_cache_path(text, backend=backend, model=model, cache_dir=cache_dir)
if cache and path.is_file():
try:
payload = json.loads(path.read_text())
return EnhanceResult(
original=str(payload.get("original", text)),
enhanced=str(payload["enhanced"]),
# Mark cache hits explicitly so metrics/logs can distinguish
# a free replay from a fresh template/mlx-lm call.
backend="cache",
elapsed_s=0.0,
model=payload.get("model"),
)
except (OSError, KeyError, json.JSONDecodeError):
pass
result = enhance_prompt(
text,
backend=backend,
model=model,
system_prompt=system_prompt,
max_tokens=max_tokens,
)
if cache:
try:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(
{
"original": result.original,
"enhanced": result.enhanced,
"backend": result.backend,
"model": result.model,
},
indent=2,
))
except OSError as exc: # pragma: no cover - cache is best-effort
logger.info("[MLX enhance] cache write skipped: %s", exc)
return result
def enhance_result_as_metrics(result: EnhanceResult | None) -> dict[str, Any]:
"""
Convert prompt enhancement results into metrics fields.
Parameters:
result (EnhanceResult | None): The enhancement result, or `None` when no enhancement was performed.
Returns:
dict[str, Any]: A metrics mapping containing enhancement status, backend metadata, timing, and original and enhanced prompts.
"""
if result is None:
return {
"enhance_prompt": False,
"enhance_backend": None,
"enhance_model": None,
"enhance_elapsed_s": None,
"prompt_original": None,
"prompt_enhanced": None,
}
return {
"enhance_prompt": True,
"enhance_backend": result.backend,
"enhance_model": result.model,
"enhance_elapsed_s": result.elapsed_s,
"prompt_original": result.original,
"prompt_enhanced": result.enhanced,
}
__all__ = [
"DEFAULT_ENHANCE_SYSTEM_PROMPT",
"DEFAULT_MLX_LM_MODEL",
"EnhanceResult",
"enhance_cache_path",
"enhance_prompt",
"enhance_prompt_mlx_lm",
"enhance_prompt_template",
"enhance_result_as_metrics",
"load_or_enhance_prompt",
]
+284
View File
@@ -0,0 +1,284 @@
# SPDX-License-Identifier: Apache-2.0
"""MLX block-scaled quantization backends (affine INT8, MXFP8/4, NVFP4).
Isolated experiment module: probes which ``mx.quantize`` modes the installed
MLX build supports and exposes a thin wrapper around native quantized matmul.
Depends only on ``mlx.core`` and the standard library — do not import the rest
of FastVideo from here.
"""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from typing import Final
from collections.abc import Mapping
import mlx.core as mx
# Probe matrix side length: must be divisible by every mode's group size
# (affine g64, mxfp* g32, nvfp4 g16).
_PROBE_DIM: Final[int] = 64
class QuantBackend(str, Enum):
"""Named MLX quantization backends evaluated for M5 Neural Accelerators."""
AFFINE_INT8_G64 = "affine_int8_g64"
MXFP8 = "mxfp8"
MXFP4 = "mxfp4"
NVFP4 = "nvfp4"
BACKENDS: Final[tuple[str, ...]] = tuple(b.value for b in QuantBackend)
# Backend name -> kwargs for mx.quantize / mx.quantized_matmul.
# Affine baseline matches FastVideo DiT load path (INT8, group size 64).
# MX/NV block-scaled modes use MLX defaults (see mx.quantize docs).
_BACKEND_KWARGS: Final[Mapping[str, Mapping[str, object]]] = {
QuantBackend.AFFINE_INT8_G64.value: {
"mode": "affine",
"bits": 8,
"group_size": 64,
},
QuantBackend.MXFP8.value: {
"mode": "mxfp8",
"bits": None,
"group_size": None,
},
QuantBackend.MXFP4.value: {
"mode": "mxfp4",
"bits": None,
"group_size": None,
},
QuantBackend.NVFP4.value: {
"mode": "nvfp4",
"bits": None,
"group_size": None,
},
}
_SUPPORT_CACHE: dict[str, bool] = {}
_SUPPORT_ERROR_CACHE: dict[str, str | None] = {}
_BYTES_CACHE: dict[str, float] = {}
@dataclass(frozen=True)
class QuantizedWeight:
"""Packed quantized weight plus scales/biases for one backend."""
weight: mx.array
scales: mx.array
biases: mx.array | None
backend: str
mode: str
bits: int | None
group_size: int | None
# Original (rows, cols) of the fp weight, used for bytes-per-element.
orig_shape: tuple[int, int]
def _normalize_backend(backend: str) -> str:
"""Normalize a quantization backend name and validate that it is supported.
Parameters:
backend (str): Backend name to normalize.
Returns:
str: The lowercase backend name without surrounding whitespace.
Raises:
ValueError: If the backend name is unknown.
"""
name = backend.strip().lower()
if name not in _BACKEND_KWARGS:
known = ", ".join(BACKENDS)
raise ValueError(f"Unknown quant backend {backend!r}. Expected one of: {known}")
return name
def _kwargs_for(backend: str) -> dict[str, object]:
"""Return the MLX quantization arguments configured for a backend.
Parameters:
backend (str): Backend name to resolve.
Returns:
dict[str, object]: Quantization arguments for the normalized backend.
"""
return dict(_BACKEND_KWARGS[_normalize_backend(backend)])
def support_error(backend: str) -> str | None:
"""
Check whether a quantization backend is supported by the current MLX runtime.
Parameters:
backend (str): Quantization backend name.
Returns:
str | None: An error description when the backend is unsupported, or `None` when supported.
"""
name = _normalize_backend(backend)
if name in _SUPPORT_ERROR_CACHE:
return _SUPPORT_ERROR_CACHE[name]
kwargs = _kwargs_for(name)
try:
w = mx.zeros((_PROBE_DIM, _PROBE_DIM), dtype=mx.float16)
quantized = mx.quantize(
w,
group_size=kwargs["group_size"], # type: ignore[arg-type]
bits=kwargs["bits"], # type: ignore[arg-type]
mode=str(kwargs["mode"]),
)
w_q = quantized[0]
scales = quantized[1]
biases = quantized[2] if len(quantized) == 3 else None
x = mx.zeros((1, _PROBE_DIM), dtype=mx.float16)
y = mx.quantized_matmul(
x,
w_q,
scales,
biases,
transpose=True,
group_size=kwargs["group_size"], # type: ignore[arg-type]
bits=kwargs["bits"], # type: ignore[arg-type]
mode=str(kwargs["mode"]),
)
mx.eval(y)
_SUPPORT_ERROR_CACHE[name] = None
_SUPPORT_CACHE[name] = True
except Exception as exc: # noqa: BLE001 - MLX raises varied types per mode/version.
msg = f"{type(exc).__name__}: {exc}"
_SUPPORT_ERROR_CACHE[name] = msg
_SUPPORT_CACHE[name] = False
return _SUPPORT_ERROR_CACHE[name]
def is_supported(backend: str) -> bool:
"""Return True if the installed MLX build can quantize/matmul with ``backend``."""
name = _normalize_backend(backend)
if name not in _SUPPORT_CACHE:
support_error(name)
return _SUPPORT_CACHE[name]
def quantize_weight(w: mx.array, backend: str) -> QuantizedWeight:
"""
Quantize a two-dimensional weight matrix using the specified native MLX backend.
Parameters:
w (mx.array): The two-dimensional weight matrix to quantize.
backend (str): The quantization backend to use.
Returns:
QuantizedWeight: The quantized weights and associated quantization metadata.
Raises:
ValueError: If the backend is unknown, the weight is not two-dimensional,
or its last dimension is not divisible by the backend's group size.
RuntimeError: If the backend is unsupported by the installed MLX build.
"""
name = _normalize_backend(backend)
err = support_error(name)
if err is not None:
mlx_version = getattr(mx, "__version__", "unknown")
raise RuntimeError(f"Quant backend {name!r} is not supported by installed mlx "
f"({mlx_version}): {err}")
if w.ndim != 2:
raise ValueError(f"quantize_weight expects a 2D weight, got shape {tuple(w.shape)}")
rows, cols = int(w.shape[0]), int(w.shape[1])
kwargs = _kwargs_for(name)
group_size = kwargs["group_size"]
# When group_size is None, MLX applies the mode default; only check when set.
if isinstance(group_size, int) and cols % group_size != 0:
raise ValueError(f"Weight last dim {cols} must be divisible by group_size={group_size} "
f"for backend {name!r}")
quantized = mx.quantize(
w,
group_size=kwargs["group_size"], # type: ignore[arg-type]
bits=kwargs["bits"], # type: ignore[arg-type]
mode=str(kwargs["mode"]),
)
w_q = quantized[0]
scales = quantized[1]
biases = quantized[2] if len(quantized) == 3 else None
eval_args = [w_q, scales] if biases is None else [w_q, scales, biases]
mx.eval(*eval_args)
return QuantizedWeight(
weight=w_q,
scales=scales,
biases=biases,
backend=name,
mode=str(kwargs["mode"]),
bits=kwargs["bits"] if isinstance(kwargs["bits"], int) else None,
group_size=group_size if isinstance(group_size, int) else None,
orig_shape=(rows, cols),
)
def quantized_matmul(x: mx.array, qw: QuantizedWeight) -> mx.array:
"""Compute ``x @ w.T`` in the quantized domain via ``mx.quantized_matmul``."""
return mx.quantized_matmul(
x,
qw.weight,
qw.scales,
qw.biases,
transpose=True,
group_size=qw.group_size,
bits=qw.bits,
mode=qw.mode,
)
def _artifact_nbytes(qw: QuantizedWeight) -> int:
"""
Calculate the total storage size of a quantized weight artifact in bytes.
Parameters:
qw (QuantizedWeight): Quantized weight artifact whose packed weights, scales, and optional biases are measured.
Returns:
int: Total number of bytes used by the artifact's stored arrays.
"""
total = int(qw.weight.nbytes) + int(qw.scales.nbytes)
if qw.biases is not None:
total += int(qw.biases.nbytes)
return total
def bytes_per_weight(backend: str) -> float:
"""
Measure the effective storage cost of a quantized weight.
Parameters:
backend (str): Quantization backend to measure.
Returns:
float: Stored bytes per original weight element, including packed weights,
scales, and optional biases.
Raises:
RuntimeError: If the backend is unsupported.
"""
name = _normalize_backend(backend)
if name in _BYTES_CACHE:
return _BYTES_CACHE[name]
err = support_error(name)
if err is not None:
mlx_version = getattr(mx, "__version__", "unknown")
raise RuntimeError(f"Cannot measure bytes_per_weight for unsupported backend {name!r} "
f"(mlx {mlx_version}): {err}")
probe = mx.zeros((_PROBE_DIM, _PROBE_DIM), dtype=mx.float16)
qw = quantize_weight(probe, name)
n_elem = qw.orig_shape[0] * qw.orig_shape[1]
value = _artifact_nbytes(qw) / float(n_elem)
_BYTES_CACHE[name] = value
return value
+690
View File
@@ -0,0 +1,690 @@
# SPDX-License-Identifier: Apache-2.0
"""Two-pass spatial refine for the MLX Wan runtime (H3 / LTX-2 pattern).
Biggest quality lever on Apple Silicon without a new model or training:
generate at base resolution, then run a second denoising pass with the
*same* DiT at a higher resolution.
This is the MLX-side port of the CUDA refine template in
``fastvideo/pipelines/basic/ltx2/stages/ltx2_refine.py`` and the H3
"base + regenerate" pattern documented in
``docs/design/mac_qad_two_product_strategy.md``:
1. :func:`plan_refine_resolutions` — split the request into stage-1
(base) and stage-2 (target) pixel sizes, validating VAE / patch
alignment the way :class:`LTX2RefineInitStage` does.
2. :func:`upsample_latents_spatial` — 2× (or N×) spatial upsample of
clean latents. Wan has no learned latent upsampler on Mac, so this
is bilinear over the H×W plane (temporal axis untouched) — same
role as LTX-2's ``upsample_video`` hand-off, without the learned
residual.
3. :func:`prepare_refine_latents` — upsample + re-noise the clean
stage-1 latents to the stage-2 sigma so the second denoise has
something to refine (mirrors :class:`LTX2UpsampleStage` +
``apply_ltx2_gaussian_noiser``).
4. :func:`run_two_pass_dmd` — orchestrate stage-1 denoise → refine
hand-off → stage-2 denoise with the same model / prompt embeds.
No LoRA swap, no dedicated SR weights, no new training — pure pipeline
work reusable by Wan2.1-14B and Wan2.2-5B on Apple Silicon.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from collections.abc import Callable, Sequence
import numpy as np
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, add_noise, dmd_step
if TYPE_CHECKING: # pragma: no cover - typing only
import mlx.core as mx
logger = init_logger(__name__)
# Default stage-2 noise level when the caller does not supply a schedule.
# Matches the first entry of LTX-2's STAGE_2_DISTILLED_SIGMA_VALUES in spirit
# (start the refine denoise from a high-noise level) without hard-wiring the
# LTX-2 distilled grid onto Wan's flow-match schedule.
DEFAULT_REFINE_SIGMA = 0.909375
@dataclass(frozen=True)
class RefinePlan:
"""Resolved stage-1 / stage-2 geometry for a two-pass refine run."""
target_height: int
target_width: int
stage1_height: int
stage1_width: int
spatial_scale: int
vae_spatial_compression: int
vae_temporal_compression: int
num_frames: int
@property
def stage1_latent_height(self) -> int:
"""Return the stage-1 latent height after VAE spatial compression."""
return self.stage1_height // self.vae_spatial_compression
@property
def stage1_latent_width(self) -> int:
"""Return the stage-one latent width after VAE spatial compression."""
return self.stage1_width // self.vae_spatial_compression
@property
def stage2_latent_height(self) -> int:
"""Calculate the target-resolution latent height.
Returns:
int: The target height divided by the VAE spatial compression factor.
"""
return self.target_height // self.vae_spatial_compression
@property
def stage2_latent_width(self) -> int:
"""Return the target image width in latent-space units."""
return self.target_width // self.vae_spatial_compression
@property
def latent_frames(self) -> int:
"""Calculate the number of latent frames after VAE temporal compression.
Returns:
int: The compressed latent frame count.
"""
return (self.num_frames - 1) // self.vae_temporal_compression + 1
def plan_refine_resolutions(
*,
height: int,
width: int,
num_frames: int,
spatial_scale: int = 2,
vae_spatial_compression: int = 8,
vae_temporal_compression: int = 4,
patch_size: tuple[int, int, int] = (1, 2, 2),
enabled: bool = True,
mode_label: str = "Refine",
) -> RefinePlan:
"""
Validate the requested dimensions and create the stage-1 and target-resolution refinement plan.
Parameters:
height (int): Target image height in pixels.
width (int): Target image width in pixels.
num_frames (int): Number of frames in the input sequence.
spatial_scale (int): Factor used to reduce spatial dimensions for stage 1.
vae_spatial_compression (int): Spatial compression factor of the VAE.
vae_temporal_compression (int): Temporal compression factor of the VAE.
patch_size (tuple[int, int, int]): Temporal and spatial patch dimensions used to validate latent-grid alignment.
enabled (bool): Whether to use two-pass refinement.
mode_label (str): Name of the calling mode, used to prefix validation
errors so ``--fast-spatial`` failures do not read as refine failures.
Returns:
RefinePlan: The validated stage-1 and target-resolution plan.
"""
if height <= 0 or width <= 0:
raise ValueError(f"height/width must be positive, got {height}x{width}")
if spatial_scale < 1:
raise ValueError(f"spatial_scale must be >= 1, got {spatial_scale}")
if num_frames <= 0:
raise ValueError(f"num_frames must be positive, got {num_frames}")
if vae_spatial_compression < 1 or vae_temporal_compression < 1:
raise ValueError("VAE compression factors must be positive")
if height % vae_spatial_compression != 0 or width % vae_spatial_compression != 0:
raise ValueError(f"height/width must be divisible by vae_spatial_compression={vae_spatial_compression} "
f"(got {height}x{width}).")
if (num_frames - 1) % vae_temporal_compression != 0:
raise ValueError(f"num_frames must be 1 modulo vae_temporal_compression={vae_temporal_compression} "
f"(got {num_frames}).")
if not enabled or spatial_scale == 1:
plan = RefinePlan(
target_height=height,
target_width=width,
stage1_height=height,
stage1_width=width,
spatial_scale=1,
vae_spatial_compression=vae_spatial_compression,
vae_temporal_compression=vae_temporal_compression,
num_frames=num_frames,
)
_validate_plan(plan, patch_size=patch_size, mode_label=mode_label)
return plan
if height % spatial_scale != 0 or width % spatial_scale != 0:
raise ValueError(f"{mode_label} requires height/width divisible by spatial_scale={spatial_scale} "
f"(got {height}x{width}).")
stage1_height = height // spatial_scale
stage1_width = width // spatial_scale
# Stage-1 must land on a VAE-aligned grid so the first denoise produces
# valid latents; the LTX-2 init stage enforces the same constraint.
if (stage1_height % vae_spatial_compression != 0 or stage1_width % vae_spatial_compression != 0):
raise ValueError(f"{mode_label} requires height/width divisible by "
f"{spatial_scale * vae_spatial_compression} "
f"(got {height}x{width}, vae_spatial={vae_spatial_compression}).")
plan = RefinePlan(
target_height=height,
target_width=width,
stage1_height=stage1_height,
stage1_width=stage1_width,
spatial_scale=spatial_scale,
vae_spatial_compression=vae_spatial_compression,
vae_temporal_compression=vae_temporal_compression,
num_frames=num_frames,
)
_validate_plan(plan, patch_size=patch_size, mode_label=mode_label)
logger.info(
"[MLX refine] enabled: stage1=%dx%d stage2=%dx%d scale=%dx",
stage1_width,
stage1_height,
width,
height,
spatial_scale,
)
return plan
def _validate_plan(plan: RefinePlan, *, patch_size: tuple[int, int, int], mode_label: str = "Refine") -> None:
"""
Validate that both refinement stages have latent dimensions aligned to the patch grid.
Parameters:
patch_size (tuple[int, int, int]): Temporal, height, and width patch dimensions.
Raises:
ValueError: If a stage's spatial latent dimensions or the temporal latent
dimension is not divisible by the corresponding patch dimension.
"""
pt, ph, pw = patch_size
for label, lh, lw in (
("stage1", plan.stage1_latent_height, plan.stage1_latent_width),
("stage2", plan.stage2_latent_height, plan.stage2_latent_width),
):
if lh % ph != 0 or lw % pw != 0:
raise ValueError(f"{mode_label} {label} latent grid {lh}x{lw} is not divisible by "
f"patch spatial size {ph}x{pw}.")
if plan.latent_frames % pt != 0:
raise ValueError(f"{mode_label} latent_frames={plan.latent_frames} is not divisible by "
f"patch temporal size {pt}.")
def upsample_latents_spatial(
latents: Any,
*,
scale: int = 2,
mode: str = "bilinear",
) -> Any:
"""
Upsample the spatial dimensions of 5-D latent arrays while preserving the batch, channel, and temporal dimensions.
Parameters:
latents (Any): Latents with shape ``(B, C, T, H, W)``.
scale (int): Integer factor for enlarging the spatial dimensions.
mode (str): Interpolation mode, either ``"nearest"`` or ``"bilinear"``.
Returns:
Any: Latents with shape ``(B, C, T, H * scale, W * scale)``.
"""
if scale < 1:
raise ValueError(f"scale must be >= 1, got {scale}")
if scale == 1:
return latents
# Accept both mx.array and np.ndarray so unit tests can run without MLX.
is_mlx = hasattr(latents, "dtype") and type(latents).__module__.startswith("mlx")
if is_mlx:
return _upsample_latents_mlx(latents, scale=scale, mode=mode)
return _upsample_latents_numpy(np.asarray(latents), scale=scale, mode=mode)
def _upsample_latents_numpy(
latents: np.ndarray,
*,
scale: int,
mode: str,
) -> np.ndarray:
"""Upsample 5-D latent arrays spatially using nearest-neighbor or bilinear interpolation.
Parameters:
latents (np.ndarray): Latents with shape ``(B, C, T, H, W)``.
scale (int): Spatial upsampling factor.
mode (str): Interpolation mode, either ``"nearest"`` or ``"bilinear"``.
Returns:
np.ndarray: Spatially upsampled latents with preserved batch, channel, and temporal dimensions.
Raises:
ValueError: If the latents are not five-dimensional or the interpolation mode is unsupported.
"""
if latents.ndim != 5:
raise ValueError(f"Expected 5-D latents (B,C,T,H,W), got shape {latents.shape}")
b, c, t, h, w = latents.shape
if mode == "nearest":
# (B,C,T,H,1,W,1) -> broadcast to (B,C,T,H,scale,W,scale) -> merge.
out = np.repeat(np.repeat(latents, scale, axis=3), scale, axis=4)
return out
if mode != "bilinear":
raise ValueError(f"Unsupported upsample mode: {mode}")
# Bilinear over the spatial plane. Flatten (B,C,T) into a batch of 2-D
# maps so a single vectorized gather covers every frame/channel.
src = latents.reshape(b * c * t, h, w).astype(np.float32, copy=False)
out_h, out_w = h * scale, w * scale
# Map output pixel centers onto the input grid (align_corners=False).
ys = (np.arange(out_h, dtype=np.float32) + 0.5) * (h / out_h) - 0.5
xs = (np.arange(out_w, dtype=np.float32) + 0.5) * (w / out_w) - 0.5
ys = np.clip(ys, 0.0, h - 1.0)
xs = np.clip(xs, 0.0, w - 1.0)
y0 = np.floor(ys).astype(np.int64)
x0 = np.floor(xs).astype(np.int64)
y1 = np.minimum(y0 + 1, h - 1)
x1 = np.minimum(x0 + 1, w - 1)
wy = (ys - y0.astype(np.float32))[:, None]
wx = (xs - x0.astype(np.float32))[None, :]
# Gather the four corners: shape (N, out_h, out_w).
Ia = src[:, y0[:, None], x0[None, :]]
Ib = src[:, y0[:, None], x1[None, :]]
Ic = src[:, y1[:, None], x0[None, :]]
Id = src[:, y1[:, None], x1[None, :]]
wa = (1.0 - wy) * (1.0 - wx)
wb = (1.0 - wy) * wx
wc = wy * (1.0 - wx)
wd = wy * wx
out = wa * Ia + wb * Ib + wc * Ic + wd * Id
return out.reshape(b, c, t, out_h, out_w).astype(latents.dtype, copy=False)
def _upsample_latents_mlx(latents: mx.array, *, scale: int, mode: str) -> mx.array:
"""
Upsample MLX latent tensors along their spatial dimensions.
Parameters:
latents (mx.array): A latent tensor with shape `(B, C, T, H, W)`.
scale (int): The integer spatial upsampling factor.
mode (str): The interpolation mode, such as `"nearest"` or `"bilinear"`.
Returns:
mx.array: The spatially upsampled latent tensor with its original data type.
"""
import mlx.core as mx
# Route through NumPy for the interpolation math. Latent tensors at Mac
# resolutions are small (e.g. 1×16×21×30×52 ≈ 1 MB) so the host hop is
# cheaper than carrying a bespoke Metal bilinear kernel, and it keeps
# the CPU-only unit tests and the MLX path on one implementation.
np_latents = np.array(latents.astype(mx.float32))
up = _upsample_latents_numpy(np_latents, scale=scale, mode=mode)
return mx.array(up).astype(latents.dtype)
def prepare_refine_latents(
clean_latents: Any,
*,
scale: int = 2,
sigma: float = DEFAULT_REFINE_SIGMA,
noise: Any | None = None,
add_noise_flag: bool = True,
upsample_mode: str = "bilinear",
seed: int | None = None,
) -> Any:
"""
Upsample clean latents spatially and optionally mix them with Gaussian noise.
Parameters:
clean_latents: The stage-1 latent tensor.
sigma: Noise mixing factor between 0 and 1.
noise: Optional noise tensor to mix with the upsampled latents.
add_noise_flag: Whether to apply noise mixing.
upsample_mode: Spatial interpolation mode.
seed: Optional seed for generated noise.
Returns:
The upsampled latents, optionally mixed with noise.
Raises:
ValueError: If sigma is outside the range from 0 to 1.
"""
if sigma < 0.0 or sigma > 1.0:
raise ValueError(f"sigma must be in [0, 1], got {sigma}")
upsampled = upsample_latents_spatial(clean_latents, scale=scale, mode=upsample_mode)
if not add_noise_flag or sigma == 0.0:
return upsampled
is_mlx = hasattr(upsampled, "dtype") and type(upsampled).__module__.startswith("mlx")
if noise is None:
noise = _draw_noise_like(upsampled, seed=seed, is_mlx=is_mlx)
return add_noise(upsampled, noise, float(sigma))
def refine_sigma_from_schedule(
schedule: MLXDMDSchedule,
timesteps: Sequence[float | int],
) -> float:
"""Derive the refinement noise level from the first refinement timestep.
Parameters:
schedule (MLXDMDSchedule): Schedule used to map timesteps to noise levels.
timesteps (Sequence[float | int]): Refinement timesteps, whose first value determines the sigma.
Returns:
float: Sigma corresponding to the first refinement timestep.
Raises:
ValueError: If `timesteps` is empty.
"""
if not timesteps:
raise ValueError("timesteps must be non-empty to derive a refine sigma")
return float(schedule.sigma_for(float(timesteps[0])))
def default_refine_timesteps(
schedule: MLXDMDSchedule,
timesteps: Sequence[float | int],
) -> list[float]:
"""Derive stage-2 timesteps from the stage-1 DMD grid.
The stage-2 pass must start *below* full noise, otherwise the hand-off
``(1 - sigma) * upsampled + sigma * noise`` weights stage 1 at zero and
the refine pass silently becomes a plain full-resolution generation at
twice the cost. FastWan's stage-1 grid opens at ``t=1000`` (``sigma``
exactly 1.0), so reusing it verbatim — which is what happens when
``--refine-dmd-denoising-steps`` is left unset — discards stage 1.
Dropping the leading full-noise entries keeps the pass on timesteps the
distilled student was actually trained on (no off-grid ``t`` the DiT has
never seen) while letting the stage-1 structure through.
Parameters:
schedule (MLXDMDSchedule): Schedule used to map timesteps to noise levels.
timesteps (Sequence[float | int]): The stage-1 DMD timestep grid.
Returns:
list[float]: The stage-1 grid with leading full-noise timesteps removed.
Raises:
ValueError: If every timestep in the grid is at full noise, leaving no
usable refine step.
"""
steps = [float(step) for step in timesteps]
first = 0
while first < len(steps) and schedule.sigma_for(steps[first]) >= 1.0:
first += 1
if first == len(steps):
raise ValueError(f"No usable refine timesteps in {steps}: every entry is at sigma >= 1 "
"(full noise), which would discard the stage-1 result. Pass "
"explicit stage-2 timesteps below the full-noise step.")
return steps[first:]
def run_dmd_loop(
*,
dit: Any,
latents: Any,
encoder_hidden_states: Any,
freqs_cis: tuple[Any, Any],
timesteps: Sequence[float | int],
schedule: MLXDMDSchedule,
mx_dtype: Any,
seed: int | None = None,
step_callback: Callable[[int, int], None] | None = None,
label: str = "denoise",
) -> Any:
"""
Denoise latents over the supplied timesteps using the DMD schedule.
Parameters:
timesteps (Sequence[float | int]): Denoising timesteps in execution order.
seed (int | None): Seed for reproducible intermediate noise generation.
step_callback (Callable[[int, int], None] | None): Callback receiving the
completed step number and total step count.
label (str): Label used for progress output when no callback is provided.
Returns:
Any: The denoised latents.
"""
import mlx.core as mx
renoise_rng = np.random.default_rng(seed) if seed is not None else None
latents_out = latents
n_steps = len(timesteps)
for step_index, timestep in enumerate(timesteps):
noise_input = latents_out
ts_val = float(timestep)
timestep_mx = mx.array([ts_val]).astype(mx.float32)
noise_pred = dit(
latents_out.astype(mx_dtype),
encoder_hidden_states,
timestep_mx,
freqs_cis,
)
noise_input_f32 = noise_input.astype(mx.float32)
pred_noise_f32 = noise_pred.astype(mx.float32)
if step_index < n_steps - 1:
next_ts: float | None = float(timesteps[step_index + 1])
if renoise_rng is not None:
renoise = mx.array(renoise_rng.standard_normal(tuple(noise_input_f32.shape)).astype(np.float32))
else:
renoise = mx.random.normal(noise_input_f32.shape).astype(mx.float32)
else:
next_ts, renoise = None, None
latents_out = dmd_step(
latents=noise_input_f32,
noise_input_latent=noise_input_f32,
pred_noise=pred_noise_f32,
schedule=schedule,
timestep=ts_val,
next_timestep=next_ts,
noise=renoise,
).astype(mx_dtype)
mx.eval(latents_out)
if step_callback is not None:
step_callback(step_index + 1, n_steps)
else:
print(f"{label} step {step_index + 1}/{n_steps} complete")
return latents_out
@dataclass(frozen=True)
class TwoPassResult:
"""Outputs of :func:`run_two_pass_dmd`."""
latents: Any
stage1_latents: Any
plan: RefinePlan
refine_sigma: float
def run_two_pass_dmd(
*,
dit: Any,
encoder_hidden_states: Any,
noise_latents_stage1: Any,
freqs_cis_stage1: tuple[Any, Any],
freqs_cis_stage2: tuple[Any, Any] | None,
plan: RefinePlan,
schedule: MLXDMDSchedule,
timesteps: Sequence[float | int],
refine_timesteps: Sequence[float | int] | None = None,
mx_dtype: Any,
seed: int = 0,
add_noise_flag: bool = True,
upsample_mode: str = "bilinear",
refine_sigma: float | None = None,
step_callback: Callable[[str, int, int], None] | None = None,
) -> TwoPassResult:
"""
Run base denoising and, when enabled, spatial refinement denoising.
Parameters:
dit: DiT callable used for both denoising passes.
encoder_hidden_states: Prompt embeddings shared across both passes.
noise_latents_stage1: Initial stage-1 noise latents.
freqs_cis_stage1: RoPE tables for the stage-1 resolution.
freqs_cis_stage2: RoPE tables for the stage-2 resolution, required when refinement is enabled.
plan: Refinement geometry and configuration.
schedule: Flow-matching schedule used by both passes.
timesteps: Stage-1 denoising timesteps.
refine_timesteps: Stage-2 denoising timesteps. Uses `timesteps` when omitted.
mx_dtype: MLX dtype used for DiT inputs and outputs.
seed: Base seed for reproducible noise generation.
add_noise_flag: Whether to add noise to the upsampled stage-1 latents.
upsample_mode: Spatial upsampling mode, either `"bilinear"` or `"nearest"`.
refine_sigma: Stage-2 starting noise level. Derived from the first refinement timestep when omitted.
step_callback: Optional callback receiving the phase name, step index, and total step count.
Returns:
TwoPassResult containing the final latents, stage-1 latents, refinement plan, and applied refinement sigma.
Raises:
ValueError: If refinement is enabled without stage-2 RoPE tables, without refinement timesteps, or if upsampled latents do not match the planned stage-2 dimensions.
"""
stage1_cb = None
stage2_cb = None
if step_callback is not None:
stage1_cb = lambda i, n: step_callback("stage1", i, n) # noqa: E731
stage2_cb = lambda i, n: step_callback("stage2", i, n) # noqa: E731
stage1_latents = run_dmd_loop(
dit=dit,
latents=noise_latents_stage1,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis_stage1,
timesteps=timesteps,
schedule=schedule,
mx_dtype=mx_dtype,
seed=seed,
step_callback=stage1_cb,
label="stage1 denoise",
)
if plan.spatial_scale == 1:
return TwoPassResult(
latents=stage1_latents,
stage1_latents=stage1_latents,
plan=plan,
refine_sigma=0.0,
)
if freqs_cis_stage2 is None:
raise ValueError("freqs_cis_stage2 is required when refine spatial_scale > 1")
if refine_timesteps is not None:
stage2_timesteps = [float(step) for step in refine_timesteps]
if not stage2_timesteps:
raise ValueError("refine_timesteps must be non-empty when refine is enabled")
else:
# Not `list(timesteps)`: the stage-1 grid opens at full noise, which
# would weight the stage-1 result at zero. See default_refine_timesteps.
stage2_timesteps = default_refine_timesteps(schedule, timesteps)
grid_sigma = refine_sigma_from_schedule(schedule, stage2_timesteps)
sigma = float(refine_sigma) if refine_sigma is not None else grid_sigma
if refine_sigma is not None and abs(sigma - grid_sigma) > 1e-6:
# The loop tells the DiT `stage2_timesteps[0]`, which implies grid_sigma.
# Overriding the hand-off noise level breaks that correspondence, so the
# model is denoising from a level it was not told about. Useful for
# exploring schedules that bottom out too high, but say so out loud.
logger.warning(
"[MLX refine] refine_sigma=%.4f overrides the schedule's %.4f for timestep %g; "
"the DiT is told a timestep that no longer matches the noise it receives.",
sigma,
grid_sigma,
stage2_timesteps[0],
)
# A hand-off at sigma >= 1 is `0 * upsampled + 1 * noise`: stage 1 is
# thrown away and refine degrades to a plain full-res run at 2x the cost.
# Fail loudly rather than silently burning the first pass.
if add_noise_flag and sigma >= 1.0:
raise ValueError(f"Refine hand-off sigma={sigma:.4f} (from stage-2 timestep "
f"{stage2_timesteps[0]:g}) discards the stage-1 result entirely: "
"the upsampled latents are weighted (1 - sigma) = 0. Start the "
"stage-2 grid below the full-noise timestep, or pass "
"add_noise_flag=False to hand off the clean upsample.")
stage2_input = prepare_refine_latents(
stage1_latents,
scale=plan.spatial_scale,
sigma=sigma,
add_noise_flag=add_noise_flag,
upsample_mode=upsample_mode,
seed=seed + 1,
)
# Shape guard: upsampled latents must match the stage-2 RoPE grid.
expected_h = plan.stage2_latent_height
expected_w = plan.stage2_latent_width
got_h, got_w = int(stage2_input.shape[-2]), int(stage2_input.shape[-1])
if got_h != expected_h or got_w != expected_w:
raise ValueError(f"Refine upsample produced {got_h}x{got_w} latents, expected "
f"{expected_h}x{expected_w} for target "
f"{plan.target_height}x{plan.target_width}.")
logger.info(
"[MLX refine] stage2 start: latent=%dx%d sigma=%.4f steps=%d",
expected_w,
expected_h,
sigma,
len(stage2_timesteps),
)
stage2_latents = run_dmd_loop(
dit=dit,
latents=stage2_input,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis_stage2,
timesteps=stage2_timesteps,
schedule=schedule,
mx_dtype=mx_dtype,
seed=seed + 2,
step_callback=stage2_cb,
label="stage2 refine",
)
return TwoPassResult(
latents=stage2_latents,
stage1_latents=stage1_latents,
plan=plan,
refine_sigma=sigma,
)
__all__ = [
"DEFAULT_REFINE_SIGMA",
"RefinePlan",
"TwoPassResult",
"default_refine_timesteps",
"plan_refine_resolutions",
"prepare_refine_latents",
"refine_sigma_from_schedule",
"run_dmd_loop",
"run_two_pass_dmd",
"upsample_latents_spatial",
]
def _draw_noise_like(like: Any, *, seed: int | None, is_mlx: bool) -> Any:
"""Generate Gaussian noise with the shape and array type of the input."""
shape = tuple(int(s) for s in like.shape)
if seed is not None:
rng = np.random.default_rng(seed)
noise_np = rng.standard_normal(shape).astype(np.float32)
if is_mlx:
import mlx.core as mx
return mx.array(noise_np).astype(mx.float32)
return noise_np.astype(np.asarray(like).dtype, copy=False)
if is_mlx:
import mlx.core as mx
return mx.random.normal(shape).astype(mx.float32)
return np.random.standard_normal(shape).astype(np.asarray(like).dtype, copy=False)
+137
View File
@@ -0,0 +1,137 @@
# SPDX-License-Identifier: Apache-2.0
"""Small MLX RIFE wrapper for frame interpolation experiments.
The backend is the Apple-Silicon-native ``rife-mlx`` package, using the
``mlx-community/RIFE-4.25`` weights. Frames are HWC RGB ``uint8`` arrays.
"""
from __future__ import annotations
from collections.abc import Iterable
from functools import lru_cache
import numpy as np
from huggingface_hub.utils import LocalEntryNotFoundError
class RIFEBackendError(RuntimeError):
"""Raised when the MLX RIFE backend cannot be loaded or run."""
class RIFEWeightsUnavailableError(RIFEBackendError):
"""Raised when uncached RIFE weights cannot be downloaded."""
def aligned_keyframe_count(target_frames: int, factor: int, temporal_compression: int = 4) -> int:
"""Return the smallest VAE-aligned keyframe count that RIFE can expand to the target."""
if target_frames < 1:
raise ValueError(f"target_frames must be >= 1, got {target_frames}")
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
if temporal_compression < 1:
raise ValueError(f"temporal_compression must be >= 1, got {temporal_compression}")
required_intervals = (target_frames - 1 + factor - 1) // factor
aligned_intervals = ((required_intervals + temporal_compression - 1) // temporal_compression * temporal_compression)
return aligned_intervals + 1
def _require_hwc_rgb(frame: np.ndarray, index: int) -> np.ndarray:
array = np.asarray(frame)
if array.ndim != 3 or array.shape[2] != 3:
raise ValueError(f"frame {index} must have shape HxWx3, got {array.shape}")
if array.dtype != np.uint8:
array = np.clip(array, 0, 255).astype(np.uint8)
return np.ascontiguousarray(array)
@lru_cache(maxsize=2)
def load_model(version: str = "4.25", weights_dir: str | None = None):
"""Load the MLX-native RIFE model.
``weights_dir`` is passed through to ``build_model`` in the vendored ``rife_mlx``.
When it is ``None``, the package downloads/uses the Hugging Face
``mlx-community/RIFE-4.25`` snapshot.
"""
try:
from fastvideo.third_party.rife_mlx.utils.weights import build_model
except ImportError:
# Fall back to a separately installed upstream package, for anyone who
# already has one in the environment.
try:
from rife_mlx.utils.weights import build_model
except ImportError as exc:
raise RIFEBackendError("MLX RIFE backend is unavailable. It ships vendored under "
"fastvideo/third_party/rife_mlx, so this usually means MLX "
"itself is missing: install with `uv pip install -e '.[mlx]'`.") from exc
try:
return build_model(version, weights_dir=weights_dir)
except LocalEntryNotFoundError as exc:
raise RIFEWeightsUnavailableError(f"MLX RIFE {version} weights are unavailable: {exc}") from exc
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
raise RIFEBackendError(f"Failed to load MLX RIFE {version}: {exc}") from exc
def interpolate_pair(
frame_a: np.ndarray,
frame_b: np.ndarray,
timestep: float = 0.5,
*,
model=None,
scale: float = 1.0,
) -> np.ndarray:
"""Interpolate one RGB frame between two input RGB frames."""
if not 0.0 < timestep < 1.0:
raise ValueError(f"timestep must be inside (0, 1), got {timestep}")
img0 = _require_hwc_rgb(frame_a, 0)
img1 = _require_hwc_rgb(frame_b, 1)
if img0.shape != img1.shape:
raise ValueError(f"frame shapes must match, got {img0.shape} and {img1.shape}")
if model is None:
model = load_model()
try:
try:
from fastvideo.third_party.rife_mlx.pipeline_mlx import interpolate_pair as _interpolate_pair
except ImportError:
from rife_mlx.pipeline_mlx import interpolate_pair as _interpolate_pair
return _interpolate_pair(model, img0, img1, timestep=timestep, scale=scale)
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
raise RIFEBackendError(f"MLX RIFE interpolation failed at timestep={timestep}: {exc}") from exc
def interpolate(
frames: list[np.ndarray] | Iterable[np.ndarray],
factor: int = 2,
*,
model=None,
scale: float = 1.0,
) -> list[np.ndarray]:
"""Return an Nx interpolated frame list.
For ``len(frames)=41`` and ``factor=2``, the output length is 81:
``(41 - 1) * 2 + 1``. Original keyframes are preserved in order and RIFE
fills ``factor - 1`` intermediate timesteps between each adjacent pair.
"""
frame_list = [_require_hwc_rgb(frame, idx) for idx, frame in enumerate(frames)]
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
if len(frame_list) < 2 or factor == 1:
return [frame.copy() for frame in frame_list]
first_shape = frame_list[0].shape
for idx, frame in enumerate(frame_list[1:], start=1):
if frame.shape != first_shape:
raise ValueError(f"all frames must have the same shape; frame 0={first_shape}, frame {idx}={frame.shape}")
if model is None:
model = load_model()
out: list[np.ndarray] = []
for left, right in zip(frame_list[:-1], frame_list[1:], strict=True):
out.append(left)
for step in range(1, factor):
out.append(interpolate_pair(left, right, step / factor, model=model, scale=scale))
out.append(frame_list[-1])
return out
+139
View File
@@ -0,0 +1,139 @@
# SPDX-License-Identifier: Apache-2.0
"""On-device (MLX) DMD sampling for the FastWan runtime.
The hybrid proof-of-concept ran the FastWan DiT in MLX but bounced every
denoising step back through torch/NumPy to run the DMD scheduler math
(``MLX -> np.array -> torch (CPU) -> np.array -> MLX``). That host round-trip
forces a full device sync per step and defeats MLX's lazy graph execution.
This module mirrors the exact DMD arithmetic from
``fastvideo/models/utils.py::pred_noise_to_pred_video`` and
``FlowMatchEulerDiscreteScheduler.add_noise`` while keeping every large tensor
on the MLX device. The schedule lookup (``argmin`` over the ~1000-entry
training schedule) is done once on the host in NumPy: it is tiny, it is the
same value torch would compute, and it sidesteps the reduction-index quirk that
affects ``argmin`` on the Metal/MPS backends (see the CPU fallbacks in
``fastvideo/models/utils.py`` and ``scheduling_flow_match_euler_discrete.py``).
Because the DMD loop applies a single scalar timestep per step, ``sigma`` is a
scalar and the update is a plain elementwise affine combination — no
permute/flatten reshaping is required.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
import numpy as np
if TYPE_CHECKING: # pragma: no cover - typing only
import mlx.core as mx
@dataclass(frozen=True)
class MLXDMDSchedule:
"""Host-side copy of a flow-match scheduler's ``(sigmas, timesteps)``.
Holds the full training schedule so a DMD timestep (e.g. one of
``1000, 757, 522``) can be mapped to its flow-match ``sigma`` with the same
nearest-timestep lookup the torch path uses.
"""
sigmas: np.ndarray
timesteps: np.ndarray
@classmethod
def from_torch_scheduler(cls, scheduler: Any) -> MLXDMDSchedule:
"""Snapshot ``scheduler.sigmas`` / ``scheduler.timesteps`` to NumPy.
Matches ``pred_noise_to_pred_video`` / ``add_noise``, which index the
scheduler's *full* training schedule (not the per-inference subset).
"""
sigmas = scheduler.sigmas.detach().to("cpu").double().numpy()
timesteps = scheduler.timesteps.detach().to("cpu").double().numpy()
return cls(sigmas=np.asarray(sigmas), timesteps=np.asarray(timesteps))
def sigma_for(self, timestep: float) -> float:
"""
Find the sigma associated with the scheduled timestep nearest to the given timestep.
Parameters:
timestep (float): Timestep for which to find the nearest scheduled sigma.
Returns:
float: Sigma associated with the nearest scheduled timestep.
"""
idx = int(np.argmin(np.abs(self.timesteps - float(timestep))))
return float(self.sigmas[idx])
def pred_noise_to_pred_video(
pred_noise: mx.array,
noise_input_latent: mx.array,
sigma: float,
) -> mx.array:
"""
Compute the clean latent prediction from a flow-matching noise prediction.
Parameters:
pred_noise (mx.array): Predicted noise.
noise_input_latent (mx.array): Noised latent input.
sigma (float): Noise level used for the prediction.
Returns:
mx.array: Predicted clean latent.
"""
return noise_input_latent - sigma * pred_noise
def add_noise(
clean_latent: mx.array,
noise: mx.array,
sigma: float,
) -> mx.array:
"""Flow-match forward noising, mirroring the scheduler's ``add_noise``.
``sample = (1 - sigma) * clean_latent + sigma * noise``.
"""
return (1.0 - sigma) * clean_latent + sigma * noise
def dmd_step(
*,
latents: mx.array,
noise_input_latent: mx.array,
pred_noise: mx.array,
schedule: MLXDMDSchedule,
timestep: float,
next_timestep: float | None,
noise: mx.array | None = None,
) -> mx.array:
"""
Compute one DMD sampling update, optionally re-noising the clean latent prediction.
Args:
latents: Retained for call-site compatibility and not used in the update.
noise_input_latent: Noisy latent used to compute the clean prediction.
pred_noise: Predicted noise or velocity.
schedule: Flow-matching schedule used to map timesteps to sigmas.
timestep: Current sampling timestep.
next_timestep: Timestep for the next update, or `None` for the final step.
noise: Fresh noise used for re-noising intermediate steps.
Returns:
The re-noised latent for the next step or the clean latent prediction on
the final step.
Raises:
ValueError: If `next_timestep` is provided without `noise`.
"""
del latents # symmetry with the torch loop; not needed for the math.
sigma = schedule.sigma_for(timestep)
pred_video = pred_noise_to_pred_video(pred_noise, noise_input_latent, sigma)
if next_timestep is None:
return pred_video
if noise is None:
raise ValueError("dmd_step requires `noise` when `next_timestep` is set (re-noise step).")
sigma_next = schedule.sigma_for(next_timestep)
return add_noise(pred_video, noise, sigma_next)
+191
View File
@@ -0,0 +1,191 @@
# SPDX-License-Identifier: Apache-2.0
"""Optional TAEHV decode helpers for Apple Silicon FastWan experiments.
The TAEHV module itself is vendored at ``fastvideo/third_party/taehv`` (MIT,
madebyollin/taehv), so no source code is downloaded or executed at runtime.
Only the ``taew2_1.pth`` checkpoint is fetched on demand, and its sha256 is
verified before use.
"""
from __future__ import annotations
import hashlib
import importlib.util
import urllib.request
from pathlib import Path
import numpy as np
TAEW2_1_CHECKPOINT_URL = "https://raw.githubusercontent.com/madebyollin/taehv/main/taew2_1.pth"
# sha256 of the upstream taew2_1.pth this module was validated against
# (fetched 2026-07-02). If upstream publishes a new checkpoint, revalidate the
# decode path and update this pin.
TAEW2_1_CHECKPOINT_SHA256 = "d26151e76cdc2c9424bef988de874b33d9a53f30ef3060cd556c429c469c797e"
# Wan2.2 5B (z_dim=48) — see madebyollin/taehv taew2_2.pth; prefer
# ``fastvideo.mlx_runtime.wan_vae.ensure_taehv_checkpoint(z_dim=48)`` for new code.
TAEW2_2_CHECKPOINT_URL = ("https://raw.githubusercontent.com/madebyollin/taehv/"
"563f40bdc820ed86bcad72ea515ee48f06bd22ec/taew2_2.pth")
def _default_cache_dir() -> Path:
"""Return the default directory used to cache TAEHV checkpoints.
Returns:
Path: The TAEHV checkpoint cache directory under the user's home directory.
"""
return Path.home() / ".cache" / "fastvideo" / "taehv"
def _sha256(path: Path) -> str:
"""
Compute the SHA-256 digest of a file.
Parameters:
path (Path): The file whose contents are hashed.
Returns:
str: The file's SHA-256 digest in hexadecimal form.
"""
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def _verify_checkpoint(path: Path) -> None:
"""Verify that a TAEW2.1 checkpoint matches the expected SHA-256 digest.
Parameters:
path (Path): Path to the checkpoint file.
Raises:
RuntimeError: If the checkpoint digest does not match the expected value.
"""
actual = _sha256(path)
if actual != TAEW2_1_CHECKPOINT_SHA256:
raise RuntimeError(f"TAEHV checkpoint at {path} failed sha256 verification "
f"(expected {TAEW2_1_CHECKPOINT_SHA256}, got {actual}). "
"Delete the file to re-download it, or pass --taehv-checkpoint-path "
"pointing at a checkpoint you trust.")
def ensure_taew2_1_checkpoint(checkpoint_path: Path | None = None) -> Path:
"""
Ensure the TAEW2.1 checkpoint is available locally.
A caller-provided path is treated as trusted and is only checked for existence.
The module-managed cached checkpoint is verified against the pinned SHA-256 digest
after downloading or before reuse.
Parameters:
checkpoint_path (Path | None): Optional path to a caller-provided checkpoint.
Returns:
Path: The available checkpoint path.
Raises:
FileNotFoundError: If a caller-provided checkpoint does not exist.
RuntimeError: If a module-managed checkpoint fails verification.
"""
if checkpoint_path is not None:
if not checkpoint_path.exists():
raise FileNotFoundError(f"TAEHV checkpoint not found: {checkpoint_path}")
return checkpoint_path
checkpoint_path = _default_cache_dir() / "taew2_1.pth"
if not checkpoint_path.exists():
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
print(f"Downloading {TAEW2_1_CHECKPOINT_URL} -> {checkpoint_path}")
import socket
import tempfile
# Download to a temporary file, verify, then atomically rename.
with tempfile.NamedTemporaryFile(
mode="wb",
dir=checkpoint_path.parent,
prefix=".tmp_taew2_1_",
suffix=".pth",
delete=False,
) as tmp_file:
tmp_path = Path(tmp_file.name)
try:
old_timeout = socket.getdefaulttimeout()
socket.setdefaulttimeout(300)
try:
urllib.request.urlretrieve(
TAEW2_1_CHECKPOINT_URL,
tmp_path, # noqa: S310 - pinned public artifact, hash-verified below.
)
finally:
socket.setdefaulttimeout(old_timeout)
_verify_checkpoint(tmp_path)
tmp_path.replace(checkpoint_path)
except Exception:
tmp_path.unlink(missing_ok=True)
raise
else:
_verify_checkpoint(checkpoint_path)
return checkpoint_path
def _load_taehv_class(source_path: Path | None):
"""Load the TAEHV class from the vendored implementation or a local source override.
Parameters:
source_path (Path | None): Path to a local Python file defining `TAEHV`; `None` selects the vendored implementation.
Returns:
The loaded `TAEHV` class.
Raises:
RuntimeError: If the specified source cannot be loaded.
"""
if source_path is None:
from fastvideo.third_party.taehv import TAEHV
return TAEHV
# Explicit local override for experimenting with a modified TAEHV; this is
# a user-supplied file on disk, never something this module downloads.
spec = importlib.util.spec_from_file_location("fastvideo_external_taehv", source_path)
if spec is None or spec.loader is None:
raise RuntimeError(f"Could not load TAEHV source from {source_path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.TAEHV
def decode_latents_to_video_taehv(
*,
latents_np: np.ndarray,
output_path: Path,
fps: int,
device,
dtype,
parallel: bool,
source_path: Path | None = None,
checkpoint_path: Path | None = None,
) -> None:
"""Decode Wan/FastWan diffusion latents with TAEW2.1 and export MP4.
TAEHV's Wan wrapper expects the diffusion latents directly, without applying
the standard Wan VAE's `latents_mean` / `latents_std` shift.
"""
import torch
from diffusers.utils import export_to_video
checkpoint_path = ensure_taew2_1_checkpoint(checkpoint_path)
TAEHV = _load_taehv_class(source_path)
taehv = TAEHV(str(checkpoint_path)).to(device=device, dtype=dtype)
taehv.eval()
latents = torch.from_numpy(latents_np).to(device=device, dtype=dtype)
with torch.no_grad():
video_ntchw = taehv.decode_video(
latents.transpose(1, 2),
parallel=parallel,
show_progress_bar=False,
)
video = video_ntchw.transpose(1, 2)
video_np = video[0].permute(1, 2, 3, 0).float().cpu().numpy()
output_path.parent.mkdir(parents=True, exist_ok=True)
export_to_video(video_np, str(output_path), fps=fps)
+286
View File
@@ -0,0 +1,286 @@
# SPDX-License-Identifier: Apache-2.0
"""Wan2.2-TI2V-5B dense MLX runtime — Track D.
The Wan2.2 TI2V-5B (FullAttn) differs from the ported Wan2.1-T2V only in:
- **Scale** (24 heads x 128, hidden 3072, ffn 14336) — pure config, block math
identical, so the dense loader ``mlx_dit_from_diffusers_safetensors`` loads the
weights unchanged and we re-wrap the blocks here.
- **Per-token timestep conditioning** (``expand_timesteps=True``): the timestep is
``[batch, seq_len]`` (a level per patch token — how TI2V keeps the conditioning
image frame at t=0 while the video frames are noised). ``timestep_proj`` becomes
``[batch, seq_len, 6, dim]`` and the block/output modulation is per-token
(``[B, L, dim]``), a direct broadcast — this module implements exactly that.
I2V rides on the same forward: encode the image, replace the first latent frame,
and set that frame's timestep to 0 (handled by the caller / sampler). See
``docs/design/ti2v_5b_port_guide.md``.
"""
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING, Any
from collections.abc import Callable
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.fastwan import (
MLXWanT2VCrossAttention,
gelu_tanh,
layer_norm,
linear,
mlx_dit_from_diffusers_safetensors,
rms_norm,
silu,
timestep_embedding,
weight_dtype,
)
if TYPE_CHECKING:
import mlx.core as mx
logger = init_logger(__name__)
class MLXWan22TransformerBlock:
"""Dense Wan block with per-token (``[B, L, dim]``) timestep modulation."""
def __init__(
self,
weights: dict[str, mx.array],
*,
dim: int,
ffn_dim: int,
num_heads: int,
eps: float = 1e-6,
):
self.weights = weights
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
self.attn2 = MLXWanT2VCrossAttention(weights, dim=dim, num_heads=num_heads, eps=eps)
def __call__(self, hidden_states, encoder_hidden_states, timestep_proj, cos, sin) -> mx.array:
import mlx.core as mx
orig_dtype = hidden_states.dtype
batch = hidden_states.shape[0]
# timestep_proj: [B, L, 6, dim] -> six per-token [B, L, dim] modulations.
e = self.weights["scale_shift_table"][None] + timestep_proj.astype(mx.float32)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = [
part.squeeze(2) for part in mx.split(e, 6, axis=2)
]
# 1. Self-attention (dense, bidirectional) with per-token modulation.
norm_hidden = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm_hidden = (norm_hidden * (1.0 + scale_msa) + shift_msa).astype(orig_dtype)
query = linear(norm_hidden, self.weights["to_q.weight"], self.weights.get("to_q.bias"))
key = linear(norm_hidden, self.weights["to_k.weight"], self.weights.get("to_k.bias"))
value = linear(norm_hidden, self.weights["to_v.weight"], self.weights.get("to_v.bias"))
query = rms_norm(query, self.weights["norm_q.weight"],
eps=self.eps).reshape(batch, -1, self.num_heads, self.head_dim)
key = rms_norm(key, self.weights["norm_k.weight"], eps=self.eps).reshape(batch, -1, self.num_heads,
self.head_dim)
value = value.reshape(batch, -1, self.num_heads, self.head_dim)
from fastvideo.mlx_runtime.fastwan import apply_rotary_emb
query = apply_rotary_emb(query, cos, sin, is_neox_style=False)
key = apply_rotary_emb(key, cos, sin, is_neox_style=False)
attn = mx.fast.scaled_dot_product_attention(
query.transpose(0, 2, 1, 3),
key.transpose(0, 2, 1, 3),
value.transpose(0, 2, 1, 3),
scale=self.head_dim**-0.5,
).transpose(0, 2, 1, 3)
attn = attn.reshape(batch, -1, self.dim)
attn = linear(attn, self.weights["to_out.weight"], self.weights.get("to_out.bias"))
hidden_states = hidden_states + (attn * gate_msa).astype(orig_dtype)
norm_hidden = layer_norm(hidden_states.astype(mx.float32),
weight=self.weights["self_attn_residual_norm.norm.weight"],
bias=self.weights["self_attn_residual_norm.norm.bias"],
eps=self.eps).astype(orig_dtype)
# 2. Cross-attention, then per-token shift/scale modulation.
cross = self.attn2(norm_hidden, encoder_hidden_states)
hidden_states = hidden_states + cross
norm_hidden = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm_hidden = (norm_hidden * (1.0 + c_scale_msa) + c_shift_msa).astype(orig_dtype)
# 3. Feed-forward with per-token gate.
ff = linear(norm_hidden, self.weights["ffn.fc_in.weight"], self.weights.get("ffn.fc_in.bias"))
ff = gelu_tanh(ff)
ff = linear(ff, self.weights["ffn.fc_out.weight"], self.weights.get("ffn.fc_out.bias"))
hidden_states = hidden_states + (ff * c_gate_msa).astype(orig_dtype)
return hidden_states.astype(orig_dtype)
class MLXWan22DiT:
"""Wan2.2-TI2V-5B dense DiT with per-token timestep conditioning."""
def __init__(
self,
weights: dict[str, mx.array],
blocks: list[MLXWan22TransformerBlock],
config: dict,
*,
compile: bool = False,
) -> None:
import os
self.weights = weights
self.blocks = blocks
self.config = config
self.num_heads = int(config["num_attention_heads"])
self.head_dim = int(config["attention_head_dim"])
self.hidden_size = self.num_heads * self.head_dim
self.freq_dim = int(config["freq_dim"])
self.patch_size = tuple(config["patch_size"])
self.out_channels = int(config["out_channels"])
self.eps = float(config.get("eps", 1e-6))
self._enable_compile = compile or os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1"
self._compiled_forward: Callable[..., Any] | None = None
self._compiled_signature: tuple | None = None
def _patch_embed(self, hidden_states) -> mx.array:
batch, channels, frames, height, width = hidden_states.shape
pt, ph, pw = self.patch_size
patch_dim = channels * pt * ph * pw
x = hidden_states.reshape(batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw)
x = x.transpose(0, 2, 4, 6, 1, 3, 5, 7).reshape(batch, -1, patch_dim)
return linear(x, self.weights["patch_embedding.weight"], self.weights.get("patch_embedding.bias"))
def _condition(self, timestep, encoder_hidden_states) -> tuple:
"""Per-token conditioning. ``timestep`` is ``[B, L]`` (one level per token)."""
batch, seq = timestep.shape
t_freq = timestep_embedding(timestep.reshape(-1), self.freq_dim).astype(
weight_dtype(self.weights["condition_embedder.time_embedder.linear_1.weight"]))
temb = linear(t_freq, self.weights["condition_embedder.time_embedder.linear_1.weight"],
self.weights["condition_embedder.time_embedder.linear_1.bias"])
temb = silu(temb)
temb = linear(temb, self.weights["condition_embedder.time_embedder.linear_2.weight"],
self.weights["condition_embedder.time_embedder.linear_2.bias"])
timestep_proj = linear(silu(temb), self.weights["condition_embedder.time_proj.weight"],
self.weights["condition_embedder.time_proj.bias"])
timestep_proj = timestep_proj.reshape(batch, seq, 6, self.hidden_size)
ehs = linear(encoder_hidden_states, self.weights["condition_embedder.text_embedder.linear_1.weight"],
self.weights["condition_embedder.text_embedder.linear_1.bias"])
ehs = gelu_tanh(ehs)
ehs = linear(ehs, self.weights["condition_embedder.text_embedder.linear_2.weight"],
self.weights["condition_embedder.text_embedder.linear_2.bias"])
temb_out = temb.reshape(batch, seq, self.hidden_size)
return temb_out, timestep_proj, ehs
def _output(self, hidden_states, temb_out, *, batch, frames, height, width) -> mx.array:
import mlx.core as mx
pt, ph, pw = self.patch_size
post_pt, post_ph, post_pw = frames // pt, height // ph, width // pw
# Per-token output modulation: scale_shift_table[1,2,dim] + temb[B,L,1,dim].
e = self.weights["scale_shift_table"][None] + temb_out[:, :, None, :].astype(mx.float32)
shift, scale = [part.squeeze(2) for part in mx.split(e, 2, axis=2)]
norm = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm = (norm * (1.0 + scale) + shift).astype(weight_dtype(self.weights["proj_out.weight"]))
out = linear(norm, self.weights["proj_out.weight"], self.weights["proj_out.bias"])
out = out.reshape(batch, post_pt, post_ph, post_pw, pt, ph, pw, self.out_channels)
out = out.transpose(0, 7, 1, 4, 2, 5, 3, 6)
return out.reshape(batch, self.out_channels, frames, height, width)
def _forward(self, hidden_states, encoder_hidden_states, timestep, cos, sin) -> mx.array:
batch, _, frames, height, width = hidden_states.shape
hidden = self._patch_embed(hidden_states)
temb_out, timestep_proj, ehs = self._condition(timestep, encoder_hidden_states)
for block in self.blocks:
hidden = block(hidden, ehs, timestep_proj, cos, sin)
return self._output(hidden, temb_out, batch=batch, frames=frames, height=height, width=width)
def __call__(self, hidden_states, encoder_hidden_states, timestep, freqs_cis) -> mx.array:
cos, sin = freqs_cis
if self._enable_compile and cos is not None:
import mlx.core as mx
# One traced graph per input signature, and each pins its own copy of
# the quantized weights. --refine denoises at two resolutions, so
# keeping both alive doubles resident DiT memory. Retire the previous
# graph when the signature changes.
signature = (hidden_states.shape, encoder_hidden_states.shape, timestep.shape)
if self._compiled_forward is not None and signature != self._compiled_signature:
self._compiled_forward = None
self._compiled_signature = None
mx.clear_cache()
if self._compiled_forward is None:
self._compiled_forward = mx.compile(self._forward)
self._compiled_signature = signature
compiled_forward = self._compiled_forward
try:
return compiled_forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
except Exception as exc: # noqa: BLE001 - some graphs may not trace; fall back to eager.
logger.warning(
"Wan2.2 mx.compile forward failed (%s); falling back to eager execution.",
exc,
)
self._enable_compile = False
self._compiled_forward = None
self._compiled_signature = None
return self._forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
def mlx_wan22_dit_from_diffusers_safetensors(
checkpoint_path: str | Path,
config_path: str | Path,
*,
dtype: str = "fp16",
num_blocks: int | None = None,
quantization=None,
compile: bool = False,
) -> MLXWan22DiT:
"""Load Wan2.2-TI2V-5B (FullAttn) into ``MLXWan22DiT`` via the dense loader."""
dense = mlx_dit_from_diffusers_safetensors(checkpoint_path,
config_path,
dtype=dtype,
num_blocks=num_blocks,
quantization=quantization)
inner_dim = int(dense.config["num_attention_heads"]) * int(dense.config["attention_head_dim"])
blocks = [
MLXWan22TransformerBlock(block.weights,
dim=inner_dim,
ffn_dim=int(dense.config["ffn_dim"]),
num_heads=int(dense.config["num_attention_heads"]),
eps=float(dense.config.get("eps", 1e-6))) for block in dense.blocks
]
return MLXWan22DiT(dense.weights, blocks, dense.config, compile=compile)
def mlx_wan22_dit_from_mlx_checkpoint(
checkpoint_dir: str | Path,
*,
compile: bool = False,
) -> MLXWan22DiT:
"""Rewrap a persisted MLX DiT checkpoint with Wan2.2 conditioning.
The generic checkpoint loader intentionally rebuilds ``MLXWanDiT`` because
it is also used by the Wan2.1 runtime. Wan2.2 TI2V has the same weight
layout but needs per-token timestep modulation, so callers must rewrap the
loaded weights and blocks as :class:`MLXWan22DiT` before sampling.
"""
from fastvideo.mlx_runtime.checkpoint import load_mlx_dit_checkpoint
dense = load_mlx_dit_checkpoint(checkpoint_dir)
inner_dim = int(dense.config["num_attention_heads"]) * int(dense.config["attention_head_dim"])
blocks = [
MLXWan22TransformerBlock(
block.weights,
dim=inner_dim,
ffn_dim=int(dense.config["ffn_dim"]),
num_heads=int(dense.config["num_attention_heads"]),
eps=float(dense.config.get("eps", 1e-6)),
) for block in dense.blocks
]
return MLXWan22DiT(dense.weights, blocks, dense.config, compile=compile)
+113
View File
@@ -0,0 +1,113 @@
# SPDX-License-Identifier: Apache-2.0
"""Dense DMD sampling for MLXWan22DiT (Wan2.2 per-token timestep).
Matches the FastVideo pipeline's warped DMD schedule (``warp_denoising_step=True``,
``dmd_denoising_steps=[1000,757,522]``, ``flow_shift=5.0`` for TI2V-5B) rather
than treating raw step indices as continuous timesteps (a bug in early demos).
"""
from __future__ import annotations
from typing import TYPE_CHECKING
from collections.abc import Sequence
import numpy as np
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step, pred_noise_to_pred_video
if TYPE_CHECKING:
import mlx.core as mx
from fastvideo.mlx_runtime.wan22 import MLXWan22DiT
def build_wan22_dmd_schedule(
dmd_denoising_steps: Sequence[int] | None = None,
*,
flow_shift: float = 5.0,
warp_denoising_step: bool = True,
) -> tuple[MLXDMDSchedule, list[float]]:
"""
Build the flow-matching schedule and continuous timesteps used for Wan2.2 DMD sampling.
Parameters:
dmd_denoising_steps (Sequence[int] | None): Denoising step values to use; defaults to 1000, 757, and 522.
flow_shift (float): Flow-matching shift applied when constructing the schedule.
warp_denoising_step (bool): Whether to convert denoising steps to scheduler-warped continuous timesteps.
Returns:
tuple[MLXDMDSchedule, list[float]]: The DMD schedule and corresponding continuous timesteps.
"""
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
steps = list(dmd_denoising_steps or [1000, 757, 522])
scheduler = FlowMatchEulerDiscreteScheduler(shift=flow_shift)
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
step_idx = torch.tensor(steps, dtype=torch.long)
if warp_denoising_step:
warped = torch.cat((scheduler.timesteps.cpu(), torch.tensor([0.0], dtype=torch.float32)))
timesteps = [float(t) for t in warped[1000 - step_idx]]
else:
timesteps = [float(s) for s in steps]
return schedule, timesteps
def sample_wan22_dmd(
model: MLXWan22DiT,
encoder_hidden_states: mx.array,
noise_latents: mx.array,
freqs_cis: tuple,
*,
dmd_denoising_steps: Sequence[int] | None = None,
flow_shift: float = 5.0,
warp_denoising_step: bool = True,
seed: int = 0,
) -> mx.array:
"""
Generate clean video latents from noisy latents using iterative DMD denoising.
Parameters:
noise_latents (mx.array): Initial noisy video latents.
freqs_cis (tuple): Rotary positional frequency tensors used by the model.
dmd_denoising_steps (Sequence[int] | None): DMD denoising steps, or the default schedule when omitted.
flow_shift (float): Flow-matching schedule shift.
warp_denoising_step (bool): Whether to warp the denoising timesteps.
seed (int): Seed for reproducible intermediate re-noising.
Returns:
mx.array: Denoised video latents.
"""
import mlx.core as mx
schedule, timesteps = build_wan22_dmd_schedule(dmd_denoising_steps,
flow_shift=flow_shift,
warp_denoising_step=warp_denoising_step)
# NumPy RNG so re-noise is bit-reproducible across MLX / torch A/B dumps.
renoise_rng = np.random.default_rng(seed)
latents = noise_latents
batch, _c, frames, height, width = latents.shape
pt, ph, pw = model.patch_size
tokens = (frames // pt) * (height // ph) * (width // pw)
last = len(timesteps) - 1
for i, t in enumerate(timesteps):
ts = mx.full((batch, tokens), float(t), dtype=mx.float32)
pred = model(latents.astype(mx.float16), encoder_hidden_states, ts, freqs_cis)
ni = latents.astype(mx.float32)
pn = pred.astype(mx.float32)
if i < last:
renoise = mx.array(renoise_rng.standard_normal(tuple(latents.shape)).astype(np.float32))
latents = dmd_step(
latents=ni,
noise_input_latent=ni,
pred_noise=pn,
schedule=schedule,
timestep=float(t),
next_timestep=float(timesteps[i + 1]),
noise=renoise,
).astype(latents.dtype)
else:
latents = pred_noise_to_pred_video(pn, ni, schedule.sigma_for(float(t))).astype(latents.dtype)
mx.eval(latents)
return latents
+557
View File
@@ -0,0 +1,557 @@
# SPDX-License-Identifier: Apache-2.0
"""Wan VAE decode helpers for Apple Silicon MLX inference.
Two decode backends:
1. **TAEHV (primary / fast)** — Tiny AutoEncoder (madebyollin/taehv). Fully
MLX-native Conv2d path. ``taew2_1.pth`` for Wan2.1 (z_dim=16),
``taew2_2.pth`` for Wan2.2 5B (z_dim=48, patch_size=2). Expected decode
wall-clock ~seconds vs ~minutes for the full 3D VAE on MPS.
2. **Full AutoencoderKLWan (reference / quality)** — denormalize with
``latents_mean`` / ``latents_std`` then torch decode (MPS preferred). Used
for parity gates and when TAEHV is unavailable. A pure-MLX 3D-conv port of
the residual Wan2.2 decoder is left as follow-up (causal feat-cache +
residual up blocks are large); TAEHV covers the product latency path.
Diffusion latents from the DiT are **not** mean/std-normalized for TAEHV
(matching ``taehv_decode.py``); full VAE decode **does** denormalize first
(matching ``mlx_wan_prompt_to_video.decode_latents_to_video``).
"""
from __future__ import annotations
import hashlib
import json
import urllib.request
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Literal
import numpy as np
from fastvideo.mlx_runtime.memory import cleanup_mlx, cleanup_torch_mps
GIB = 1024**3
TAEW2_1_URL = "https://raw.githubusercontent.com/madebyollin/taehv/main/taew2_1.pth"
TAEW2_2_URL = ("https://raw.githubusercontent.com/madebyollin/taehv/"
"563f40bdc820ed86bcad72ea515ee48f06bd22ec/taew2_2.pth")
# Validated 2026-07-02 / 2026-07-09 against upstream madebyollin/taehv.
TAEW2_1_SHA256 = "d26151e76cdc2c9424bef988de874b33d9a53f30ef3060cd556c429c469c797e"
TAEW2_2_SHA256 = "d053e216ca50e2bb837bbcd79b85f0366bea00e5938025572382a773b74c559a"
DecodeBackend = Literal["taehv", "taehv-torch", "wan-vae"]
def _cache_dir() -> Path:
"""Return the local directory used to cache TAEHV files."""
return Path.home() / ".cache" / "fastvideo" / "taehv"
def _sha256(path: Path) -> str:
"""Compute the SHA-256 digest of a file.
Parameters:
path (Path): Path to the file to hash.
Returns:
str: Lowercase hexadecimal SHA-256 digest.
"""
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def _verify_checkpoint(path: Path, expected_digest: str) -> None:
"""
Verify a checkpoint's SHA-256 digest against a required expected digest.
Parameters:
path (Path): Path to the checkpoint file.
expected_digest (str): Expected lowercase SHA-256 digest.
Raises:
RuntimeError: If verification is enabled and the checkpoint digest does not match.
"""
if len(expected_digest) != 64 or any(char not in "0123456789abcdef" for char in expected_digest):
raise ValueError("A valid lowercase SHA-256 digest is required for bundled TAEHV checkpoints")
actual = _sha256(path)
if actual != expected_digest:
raise RuntimeError(f"TAEHV checkpoint at {path} failed sha256 verification "
f"(expected {expected_digest}, got {actual}). "
"Delete the file to re-download it.")
def ensure_taehv_checkpoint(*, z_dim: int, checkpoint_path: Path | None = None) -> Path:
"""
Return a validated TAEHV checkpoint for the specified latent channel count.
Parameters:
z_dim (int): Number of latent channels, supported values are 16 and 48.
checkpoint_path (Path | None): Optional existing checkpoint path to validate and use.
Returns:
Path: Path to the validated TAEHV checkpoint.
Raises:
FileNotFoundError: If the supplied checkpoint path does not exist.
ValueError: If no checkpoint is mapped to the specified latent channel count.
"""
if checkpoint_path is not None:
if not checkpoint_path.exists():
raise FileNotFoundError(f"TAEHV checkpoint not found: {checkpoint_path}")
return checkpoint_path
if z_dim == 16:
name, url, expect = "taew2_1.pth", TAEW2_1_URL, TAEW2_1_SHA256
elif z_dim == 48:
name, url, expect = "taew2_2.pth", TAEW2_2_URL, TAEW2_2_SHA256
else:
raise ValueError(f"No TAEHV checkpoint mapped for z_dim={z_dim} (supported: 16, 48)")
path = _cache_dir() / name
if not path.exists():
path.parent.mkdir(parents=True, exist_ok=True)
print(f"Downloading {url} -> {path}")
import socket
import tempfile
# Download to a temporary file, verify, then atomically rename.
with tempfile.NamedTemporaryFile(
mode="wb",
dir=path.parent,
prefix=f".tmp_{name}_",
suffix=".pth",
delete=False,
) as tmp_file:
tmp_path = Path(tmp_file.name)
try:
old_timeout = socket.getdefaulttimeout()
socket.setdefaulttimeout(300)
try:
urllib.request.urlretrieve(url, tmp_path) # noqa: S310 - public pinned artifact.
finally:
socket.setdefaulttimeout(old_timeout)
_verify_checkpoint(tmp_path, expect)
tmp_path.replace(path)
except Exception:
tmp_path.unlink(missing_ok=True)
raise
else:
_verify_checkpoint(path, expect)
return path
@dataclass(frozen=True)
class WanVAEConfigView:
"""Minimal config fields needed for denormalize + spatial scale."""
z_dim: int
latents_mean: tuple[float, ...]
latents_std: tuple[float, ...]
scale_factor_spatial: int = 8
scale_factor_temporal: int = 4
patch_size: int | None = None
vae_dir: Path | None = None
@classmethod
def from_vae_dir(cls, vae_dir: Path) -> WanVAEConfigView:
"""
Load Wan VAE configuration values from a directory.
Parameters:
vae_dir (Path): Directory containing the VAE ``config.json`` file.
Returns:
WanVAEConfigView: Configuration loaded from the VAE directory.
"""
cfg = json.loads((vae_dir / "config.json").read_text())
return cls(
z_dim=int(cfg["z_dim"]),
latents_mean=tuple(float(x) for x in cfg["latents_mean"]),
latents_std=tuple(float(x) for x in cfg["latents_std"]),
scale_factor_spatial=int(cfg.get("scale_factor_spatial", 8)),
scale_factor_temporal=int(cfg.get("scale_factor_temporal", 4)),
patch_size=cfg.get("patch_size"),
vae_dir=vae_dir,
)
def denormalize_latents_np(latents: np.ndarray, config: WanVAEConfigView) -> np.ndarray:
"""
Denormalize Wan VAE latent values using the configured means and standard deviations.
Parameters:
latents (np.ndarray): Latent values in normalized form.
config (WanVAEConfigView): Wan VAE latent statistics.
Returns:
np.ndarray: Denormalized latent values as float32.
"""
mean = np.asarray(config.latents_mean, dtype=np.float32).reshape(1, -1, 1, 1, 1)
std = np.asarray(config.latents_std, dtype=np.float32).reshape(1, -1, 1, 1, 1)
return latents.astype(np.float32) * std + mean
# ---------------------------------------------------------------------------
# MLX TAEHV decoder (Conv2d stack — primary fully-MLX product path)
# ---------------------------------------------------------------------------
def _mlx_conv2d(x: Any, weight: Any, bias: Any, *, stride: int = 1) -> Any:
import mlx.core as mx
# x: NCHW, weight: OIHW
y = mx.conv2d(x.transpose(0, 2, 3, 1), weight.transpose(0, 2, 3, 1), stride=stride, padding=1)
y = y.transpose(0, 3, 1, 2)
if bias is not None:
y = y + bias.reshape(1, -1, 1, 1)
return y
def _mlx_conv2d_1x1(x: Any, weight: Any, bias: Any = None) -> Any:
"""
Applies a 1×1 convolution to an MLX tensor in channel-first layout.
Parameters:
x (Any): Input tensor with shape [batch, channels, height, width].
weight (Any): Convolution weights.
bias (Any, optional): Optional output-channel bias.
Returns:
Any: The convolved tensor with shape [batch, output_channels, height, width].
"""
import mlx.core as mx
y = mx.conv2d(x.transpose(0, 2, 3, 1), weight.transpose(0, 2, 3, 1), stride=1, padding=0)
y = y.transpose(0, 3, 1, 2)
if bias is not None:
y = y + bias.reshape(1, -1, 1, 1)
return y
def _load_torch_state(path: Path) -> dict[str, np.ndarray]:
"""Load a PyTorch state dictionary as NumPy arrays.
Parameters:
path (Path): Path to the PyTorch checkpoint.
Returns:
dict[str, np.ndarray]: State dictionary with tensors converted to NumPy arrays.
"""
import torch
sd = torch.load(path, map_location="cpu", weights_only=True)
return {k: v.detach().float().cpu().numpy() for k, v in sd.items()}
class MLXTAEHVDecoder:
"""Minimal MLX port of TAEHV ``decoder`` (parallel-over-time MemBlocks)."""
def __init__(self, checkpoint_path: Path, *, z_dim: int) -> None:
"""Initialize the TAEHV decoder from a checkpoint for the specified latent dimensionality.
Parameters:
checkpoint_path (Path): Path to the TAEHV checkpoint.
z_dim (int): Number of latent channels, determining the decoder patch size.
"""
import mlx.core as mx
self.checkpoint_path = Path(checkpoint_path)
self.latent_channels = z_dim
# Derive patch_size from z_dim: 48 channels → patch_size=2, 16 → patch_size=1
self.patch_size = 2 if z_dim == 48 else 1
self.image_channels = 3
self.frames_to_trim = 3 # TGrow strides (1,2,2) → 2**2 - 1 for w2.1/w2.2 defaults
sd = _load_torch_state(self.checkpoint_path)
# Patch TGrow kernels like upstream TAEHV.patch_tgrow_layers.
self.weights = {k: mx.array(v) for k, v in sd.items()}
self._n_f = [256, 128, 64, 64]
def decode_ntchw(self, latents_ntchw: Any) -> Any:
"""
Decode latent video batches into clipped RGB frames.
Parameters:
latents_ntchw (Any): Latents with shape ``[N, T, C, H, W]`` and the
decoder's configured latent channel count.
Returns:
Any: Decoded frames with shape ``[N, T_out, 3, H_out, W_out]`` and values
clipped to the range ``[0, 1]``.
"""
import mlx.core as mx
x = latents_ntchw
n, t, c, h, w = x.shape
if c != self.latent_channels:
raise ValueError(f"expected C={self.latent_channels}, got {c}")
x = x.reshape(n * t, c, h, w)
x = self._run_decoder_parallel(x, n=n)
# Pixel-shuffle if patch_size > 1: (NT, 3*p*p, H, W) -> (NT, 3, H*p, W*p)
if self.patch_size > 1:
p = self.patch_size
nt, c_out, hh, ww = x.shape
x = x.reshape(nt, self.image_channels, p, p, hh, ww)
x = x.transpose(0, 1, 4, 2, 5, 3).reshape(nt, self.image_channels, hh * p, ww * p)
_, c_out, hh, ww = x.shape
t_out = x.shape[0] // n
x = x.reshape(n, t_out, c_out, hh, ww)
if self.frames_to_trim > 0 and t_out > self.frames_to_trim:
x = x[:, self.frames_to_trim:]
return mx.clip(x, 0.0, 1.0)
def _run_decoder_parallel(self, x: Any, *, n: int) -> Any:
"""
Apply the TAEHV decoder stack to flattened batch and temporal frames while preserving temporal memory.
"""
import mlx.core as mx
w = self.weights
def memblock(base: int, xx: Any, past: Any) -> Any:
"""
Apply a temporal memory block to the current and past feature tensors.
Parameters:
base (int): Decoder block index used to select the block weights.
xx (Any): Current feature tensor.
past (Any): Past feature tensor concatenated with the current features.
Returns:
Any: Activated feature tensor produced by the memory block.
"""
cat = mx.concatenate([xx, past], axis=1)
h = _mlx_conv2d(cat, w[f"decoder.{base}.conv.0.weight"], w.get(f"decoder.{base}.conv.0.bias"))
h = mx.maximum(h, 0.0)
h = _mlx_conv2d(h, w[f"decoder.{base}.conv.2.weight"], w.get(f"decoder.{base}.conv.2.bias"))
h = mx.maximum(h, 0.0)
h = _mlx_conv2d(h, w[f"decoder.{base}.conv.4.weight"], w.get(f"decoder.{base}.conv.4.bias"))
skip_key = f"decoder.{base}.skip.weight"
skip = _mlx_conv2d_1x1(xx, w[skip_key], None) if skip_key in w else xx
return mx.maximum(h + skip, 0.0)
def upsample2(xx: Any) -> Any:
"""
Upsample a four-dimensional tensor by a factor of two along its spatial dimensions.
Parameters:
xx (Any): Tensor with shape `(N, C, H, W)`.
Returns:
Any: Tensor with shape `(N, C, 2H, 2W)` containing replicated spatial values.
"""
nt, c, h, ww = xx.shape
xx = xx.reshape(nt, c, h, 1, ww, 1)
xx = mx.broadcast_to(xx, (nt, c, h, 2, ww, 2))
return xx.reshape(nt, c, h * 2, ww * 2)
def tgrow(base: int, xx: Any, stride: int) -> Any:
wt = w[f"decoder.{base}.conv.weight"]
out_ch = int(xx.shape[1]) * stride
if int(wt.shape[0]) > out_ch:
wt = wt[-out_ch:]
y = _mlx_conv2d_1x1(xx, wt, None)
if stride == 1:
return y
# TGrow.forward: (NT, C*stride, H, W) -> (NT*stride, C, H, W)
nt, c, h, ww = y.shape
c_in = c // stride
y = y.reshape(nt, stride, c_in, h, ww).transpose(0, 1, 2, 3, 4)
return y.reshape(nt * stride, c_in, h, ww)
def mem_past(xx: Any) -> Any:
"""
Build a temporal memory tensor containing a zero frame followed by the preceding frame at each time step.
Parameters:
xx (Any): Flattened batch and temporal tensor with shape ``(batch * time, channels, height, width)``.
Returns:
Any: Tensor with the same shape as ``xx`` containing the preceding frame for each temporal position.
"""
nt, c, h, ww = xx.shape
t_cur = nt // n
x_ = xx.reshape(n, t_cur, c, h, ww)
# pad one zero frame at t=0, align past[t] = x[t-1]
past = mx.concatenate([mx.zeros_like(x_[:, :1]), x_[:, :-1]], axis=1)
return past.reshape(nt, c, h, ww)
# 0 Clamp, 1 conv, 2 ReLU
x = mx.tanh(x / 3.0) * 3.0
x = _mlx_conv2d(x, w["decoder.1.weight"], w.get("decoder.1.bias"))
x = mx.maximum(x, 0.0)
for mem_idx in (3, 4, 5):
x = memblock(mem_idx, x, mem_past(x))
x = upsample2(x)
x = tgrow(7, x, 1)
x = _mlx_conv2d(x, w["decoder.8.weight"], w.get("decoder.8.bias"))
for mem_idx in (9, 10, 11):
x = memblock(mem_idx, x, mem_past(x))
x = upsample2(x)
x = tgrow(13, x, 2)
x = _mlx_conv2d(x, w["decoder.14.weight"], w.get("decoder.14.bias"))
for mem_idx in (15, 16, 17):
x = memblock(mem_idx, x, mem_past(x))
x = upsample2(x)
x = tgrow(19, x, 2)
x = _mlx_conv2d(x, w["decoder.20.weight"], w.get("decoder.20.bias"))
x = mx.maximum(x, 0.0)
x = _mlx_conv2d(x, w["decoder.22.weight"], w.get("decoder.22.bias"))
return x
def decode_latents_taehv_mlx(
latents_np: np.ndarray,
*,
z_dim: int | None = None,
checkpoint_path: Path | None = None,
) -> np.ndarray:
"""
Decode latent representations with the MLX TAEHV decoder.
Parameters:
latents_np (np.ndarray): Latents arranged as [B, C, T, H, W].
z_dim (int | None): Latent channel dimension used to select the decoder checkpoint.
checkpoint_path (Path | None): Optional path to a TAEHV checkpoint.
Returns:
np.ndarray: Decoded pixels arranged as [B, T, H, W, 3] with values in [0, 1].
Raises:
ValueError: If `latents_np` does not have five dimensions.
"""
import mlx.core as mx
if latents_np.ndim != 5:
raise ValueError(f"expected [B,C,T,H,W], got {latents_np.shape}")
c = latents_np.shape[1]
z = z_dim if z_dim is not None else c
ckpt = ensure_taehv_checkpoint(z_dim=z, checkpoint_path=checkpoint_path)
dec = MLXTAEHVDecoder(ckpt, z_dim=z)
# NTCHW
x = mx.array(latents_np.transpose(0, 2, 1, 3, 4).astype(np.float32))
out = dec.decode_ntchw(x) # N T C H W
mx.eval(out)
arr = np.array(out)
# B T H W C
return arr.transpose(0, 1, 3, 4, 2)
def decode_latents_wan_vae_torch(
latents_np: np.ndarray,
*,
vae_dir: Path,
device: str = "auto",
dtype_name: str = "fp16",
) -> np.ndarray:
"""Full AutoencoderKLWan decode on torch (MPS/CPU) with mean/std denormalize.
Returns pixels ``[B, T, H, W, 3]`` float in [0, 1].
"""
import torch
from diffusers import AutoencoderKLWan
from diffusers.video_processor import VideoProcessor
if device == "auto":
device = "mps" if torch.backends.mps.is_available() else "cpu"
dtype = torch.float16 if dtype_name == "fp16" and device == "mps" else torch.float32
config = WanVAEConfigView.from_vae_dir(vae_dir)
vae = AutoencoderKLWan.from_pretrained(vae_dir, torch_dtype=dtype, local_files_only=True).to(device)
vae.eval()
latents = torch.from_numpy(latents_np.astype(np.float32)).to(device=device, dtype=dtype)
mean = torch.tensor(config.latents_mean, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
inv_std = (1.0 / torch.tensor(config.latents_std, device=device, dtype=dtype)).view(1, -1, 1, 1, 1)
latents = latents / inv_std + mean # matches prompt_to_video path
with torch.no_grad():
video = vae.decode(latents, return_dict=False)[0]
video = VideoProcessor(vae_scale_factor=config.scale_factor_spatial).postprocess_video(video, output_type="np")
return video # [B, T, H, W, 3]
def decode_latents_to_video(
latents_np: np.ndarray,
output_path: Path,
*,
fps: int = 16,
backend: DecodeBackend = "taehv",
vae_dir: Path | None = None,
z_dim: int | None = None,
taehv_checkpoint: Path | None = None,
torch_device: str = "auto",
) -> dict[str, Any]:
"""Decode latent video frames and export them as an MP4 file.
Parameters:
latents_np (np.ndarray): Latent video representation to decode.
output_path (Path): Destination path for the MP4 file.
fps (int): Output video frame rate.
backend (DecodeBackend): Decoder backend to use.
vae_dir (Path | None): Directory containing the full Wan VAE when using
the ``wan-vae`` backend.
z_dim (int | None): Latent channel count for TAEHV decoding.
taehv_checkpoint (Path | None): Optional TAEHV checkpoint path.
torch_device (str): PyTorch device selection for PyTorch-based decoding.
Returns:
dict[str, Any]: Decode time in seconds, backend name, output path, frame
count, and video resolution.
Raises:
ValueError: If the full VAE backend lacks ``vae_dir`` or the backend is
unknown.
"""
import time
from diffusers.utils import export_to_video
t0 = time.perf_counter()
if backend in ("taehv", "taehv-torch"):
c = latents_np.shape[1] if z_dim is None else z_dim
if backend == "taehv":
video = decode_latents_taehv_mlx(latents_np, z_dim=c, checkpoint_path=taehv_checkpoint)
else:
# torch TAEHV (regression / parity reference)
import torch
from fastvideo.third_party.taehv import TAEHV
ckpt = ensure_taehv_checkpoint(z_dim=c, checkpoint_path=taehv_checkpoint)
if torch_device == "auto":
torch_device = "mps" if torch.backends.mps.is_available() else "cpu"
dtype = torch.float16 if torch_device == "mps" else torch.float32
model = TAEHV(str(ckpt)).to(device=torch_device, dtype=dtype).eval()
lat = torch.from_numpy(latents_np).to(device=torch_device, dtype=dtype)
with torch.no_grad():
out = model.decode_video(lat.transpose(1, 2), parallel=True, show_progress_bar=False)
video = out[0].permute(0, 2, 3, 1).float().cpu().numpy()[None, ...]
# out is NTCHW -> need BTHWC; decode_video returns NTCHW for batch
if video.ndim == 4:
video = video[None]
elif backend == "wan-vae":
if vae_dir is None:
raise ValueError("vae_dir required for wan-vae backend")
video = decode_latents_wan_vae_torch(latents_np, vae_dir=vae_dir, device=torch_device)
else:
raise ValueError(f"unknown backend {backend}")
decode_s = time.perf_counter() - t0
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
# export_to_video expects list/array of frames HxWxC
frames = video[0]
frames = np.clip(frames, 0.0, 1.0)
export_to_video(frames, str(output_path), fps=fps)
if backend == "taehv":
cleanup_mlx()
else:
if backend == "taehv-torch":
del model, lat, out
cleanup_torch_mps()
return {
"decode_s": decode_s,
"backend": backend,
"output_path": str(output_path),
"num_frames": int(frames.shape[0]),
"resolution": f"{frames.shape[2]}x{frames.shape[1]}" if frames.ndim == 4 else None,
}
+223
View File
@@ -0,0 +1,223 @@
# SPDX-License-Identifier: Apache-2.0
"""Chunked non-causal sliding-window self-attention for MLX scaling studies.
This module is intentionally standalone (``mlx.core`` + stdlib only) so it can
be micro-benchmarked without pulling in the DiT / FastVideo stack.
Window policy
-------------
**Symmetric** sliding window (non-causal). For query index ``i`` the allowed
key indices are:
sinks: ``j in [0, sink)`` (always visible to every query, if ``sink > 0``)
local: ``j in [max(0, i - half), min(S, i + half + 1))``
where ``half = window // 2``
so each query sees roughly ``window + 1`` local keys (plus any sinks outside
that range). This is appropriate for a dense, bidirectional DiT denoise pass.
Implementation note (FLOPs)
---------------------------
A full-size additive attention mask still materialises an ``O(S^2)`` score
matrix inside SDPA and does **not** reduce work. Instead we tile the sequence
into query blocks and run ``mx.fast.scaled_dot_product_attention`` only against
the union of keys that block needs (local slice ± sinks). That makes
per-block work ``O(chunk * (window + sink) * D)`` and total work
``O(S * (window + sink) * D)``.
"""
from __future__ import annotations
import mlx.core as mx
def _default_scale(head_dim: int, scale: float | None) -> float:
"""
Determine the attention scaling factor from an explicit value or head dimension.
Parameters:
head_dim (int): The attention head dimension used to derive the default scale.
scale (Optional[float]): An explicit scaling factor.
Returns:
float: The explicit scale converted to a float, or the reciprocal square root of `head_dim`.
Raises:
ValueError: If `scale` is not provided and `head_dim` is not positive.
"""
if scale is not None:
return float(scale)
if head_dim <= 0:
raise ValueError(f"head_dim must be positive, got {head_dim}")
return head_dim**-0.5
def _validate_qkv(q: mx.array, k: mx.array, v: mx.array) -> tuple[int, int, int, int]:
"""
Validate compatible rank-4 query, key, and value tensors.
Parameters:
q (mx.array): Query tensor shaped `(B, H, S, D)`.
k (mx.array): Key tensor with the same shape as `q`.
v (mx.array): Value tensor with the same shape as `q`.
Returns:
tuple[int, int, int, int]: Batch size, head count, sequence length, and head dimension.
Raises:
ValueError: If the tensors are not rank 4, do not have identical shapes, or have an empty sequence or head dimension.
"""
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
raise ValueError(f"q/k/v must be rank-4 (B, H, S, D); got shapes "
f"{q.shape}, {k.shape}, {v.shape}")
b, h, s, d = q.shape
if k.shape != (b, h, s, d) or v.shape != (b, h, s, d):
raise ValueError(f"q/k/v shapes must match exactly; got q={q.shape}, k={k.shape}, v={v.shape}")
if s == 0:
raise ValueError("sequence length S must be > 0")
if d == 0:
raise ValueError("head dim D must be > 0")
return b, h, s, d
def full_attention(
q: mx.array,
k: mx.array,
v: mx.array,
scale: float | None = None,
) -> mx.array:
"""
Compute dense scaled dot-product attention over the full sequence.
Parameters:
scale (float, optional): Attention scaling factor. If omitted, uses the
inverse square root of the head dimension.
Returns:
mx.array: Attention output with shape ``(B, H, S, D)``.
"""
_, _, _, d = _validate_qkv(q, k, v)
sc = _default_scale(d, scale)
return mx.fast.scaled_dot_product_attention(q, k, v, scale=sc)
def _concat_kv_slices(
k: mx.array,
v: mx.array,
ranges: list[tuple[int, int]],
) -> tuple[mx.array, mx.array]:
"""Concatenate non-overlapping ``[start, end)`` key/value slices along seq."""
if not ranges:
raise ValueError("ranges must be non-empty")
if len(ranges) == 1:
s0, e0 = ranges[0]
return k[:, :, s0:e0, :], v[:, :, s0:e0, :]
k_parts = [k[:, :, s:e, :] for s, e in ranges]
v_parts = [v[:, :, s:e, :] for s, e in ranges]
return mx.concatenate(k_parts, axis=2), mx.concatenate(v_parts, axis=2)
def _key_ranges_for_block(
qs: int,
qe: int,
seq_len: int,
half: int,
sink: int,
) -> list[tuple[int, int]]:
"""Return ordered, non-overlapping key ranges for a query block ``[qs, qe)``.
Symmetric local window over every query in the block, plus global sinks
``[0, sink)``. Overlap is merged into a single contiguous range when
possible so we avoid double-counting sink tokens.
"""
local_start = max(0, qs - half)
# Last query index is ``qe - 1``; its right edge is ``qe - 1 + half + 1 = qe + half``.
local_end = min(seq_len, qe + half)
if local_start >= local_end:
# Degenerate (should not happen for valid qs < qe); fall back to sinks only.
if sink > 0:
return [(0, min(sink, seq_len))]
raise ValueError(f"empty local key range for query block [{qs}, {qe})")
if sink <= 0:
return [(local_start, local_end)]
sink_end = min(sink, seq_len)
if local_start <= sink_end:
# Sinks abut or overlap the local window — one contiguous slice from 0.
return [(0, max(local_end, sink_end))]
# Gap between sinks and local window: two slices, concat at SDPA time.
return [(0, sink_end), (local_start, local_end)]
def windowed_attention(
q: mx.array,
k: mx.array,
v: mx.array,
window: int,
sink: int = 0,
scale: float | None = None,
*,
chunk_size: int | None = None,
) -> mx.array:
"""
Apply symmetric sliding-window self-attention with optional global sink positions.
Parameters:
q (mx.array): Query tensor shaped `(B, H, S, D)`.
k (mx.array): Key tensor shaped `(B, H, S, D)`.
v (mx.array): Value tensor shaped `(B, H, S, D)`.
window (int): Symmetric attention window width in tokens; must be at least 1.
sink (int): Number of leading key positions available to every query; must
be between 0 and the sequence length.
scale (Optional[float]): Softmax scale. Defaults to `1 / sqrt(D)`.
chunk_size (Optional[int]): Query block length used for chunked processing.
Defaults to the smaller of `window` and 512.
Returns:
mx.array: Attention output with the same shape as `q`.
Raises:
ValueError: If the inputs or attention parameters are invalid.
RuntimeError: If a query block has no available keys.
"""
_, _, seq_len, d = _validate_qkv(q, k, v)
if window < 1:
raise ValueError(f"window must be >= 1, got {window}")
if sink < 0:
raise ValueError(f"sink must be >= 0, got {sink}")
if sink > seq_len:
raise ValueError(f"sink ({sink}) cannot exceed sequence length ({seq_len})")
sc = _default_scale(d, scale)
half = window // 2
# When the requested window is at least the sequence length, every query can
# see every key under a symmetric policy — fall back to one dense SDPA.
# (Sinks are redundant once the full key set is used.)
if window >= seq_len:
return mx.fast.scaled_dot_product_attention(q, k, v, scale=sc)
chunk = min(window, 512) if chunk_size is None else int(chunk_size)
if chunk < 1:
raise ValueError(f"chunk_size must be >= 1, got {chunk}")
chunk = min(chunk, seq_len)
outputs: list[mx.array] = []
for qs in range(0, seq_len, chunk):
qe = min(seq_len, qs + chunk)
q_block = q[:, :, qs:qe, :]
ranges = _key_ranges_for_block(qs, qe, seq_len, half, sink)
k_block, v_block = _concat_kv_slices(k, v, ranges)
if k_block.shape[2] == 0:
raise RuntimeError(f"empty key set for query block [{qs}, {qe}) with window={window}, sink={sink}")
query_positions = mx.arange(qs, qe)[:, None]
key_positions = mx.array([position for start, end in ranges for position in range(start, end)])[None, :]
mask = mx.abs(query_positions - key_positions) <= half
if sink > 0:
mask = mask | (key_positions < sink)
out_block = mx.fast.scaled_dot_product_attention(q_block, k_block, v_block, scale=sc, mask=mask)
outputs.append(out_block)
return mx.concatenate(outputs, axis=2)
+137 -27
View File
@@ -10,6 +10,7 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo import envs
from fastvideo.attention import DistributedAttention
from fastvideo.attention.layer import DistributedAttention_VSA
from fastvideo.attention.selector import get_attn_backend
@@ -23,12 +24,50 @@ from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.minimax_h3_fusions import (
HAVE_TRITON,
fused_qknorm_rope,
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
minimax_h3_swiglu,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
from fastvideo.utils import get_compute_dtype
logger = init_logger(__name__)
MINIMAX_H3_MODALITY_NUM = 3
_CFG = MiniMaxH3Config()
_MINIMAX_H3_FUSION_NAMES = frozenset({"modulate", "qknorm_rope", "swiglu"})
def _enabled_minimax_h3_fusions(value: str | None = None) -> frozenset[str]:
"""Parse the independently switchable inference fusion set."""
raw = envs.FASTVIDEO_MINIMAX_H3_FUSIONS if value is None else value
normalized = raw.strip().lower()
if normalized in {"", "0", "none"}:
return frozenset()
if normalized in {"1", "all"}:
return _MINIMAX_H3_FUSION_NAMES
enabled = frozenset(item.strip() for item in normalized.split(",") if item.strip())
unknown = enabled - _MINIMAX_H3_FUSION_NAMES
if unknown:
supported = ",".join(sorted(_MINIMAX_H3_FUSION_NAMES))
raise ValueError(f"Unknown MiniMax H3 fusion(s) {sorted(unknown)}; expected a subset of {supported}.")
return enabled
def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool:
"""Triton kernels are inference-only and stay outside Dynamo capture.
The ``HAVE_TRITON`` check makes the eager fallback exact: on a CUDA build
whose Triton failed to import, an enabled fusion falls back instead of
hitting the strict wrappers' hard RuntimeError mid-forward.
"""
return (HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() and not torch.compiler.is_compiling())
class MiniMaxH3RotaryPosEmbed(nn.Module):
@@ -62,6 +101,7 @@ class MiniMaxH3FeedForward(nn.Module):
ffn_dim: int,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.fc_in = ReplicatedLinear(
@@ -78,11 +118,15 @@ class MiniMaxH3FeedForward(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.fc_out",
)
self.fuse_swiglu = fuse_swiglu
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc_in(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
if self.fuse_swiglu and _can_run_minimax_h3_fusion(hidden_states):
hidden_states = minimax_h3_swiglu(hidden_states)
else:
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
hidden_states, _ = self.fc_out(hidden_states)
return hidden_states
@@ -99,6 +143,7 @@ class MiniMaxH3Attention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...],
quant_config: QuantizationConfig | None,
prefix: str,
fuse_qknorm_rope: bool = False,
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -134,6 +179,7 @@ class MiniMaxH3Attention(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.to_out",
)
self.fuse_qknorm_rope = fuse_qknorm_rope
# VSA carries a learned gate on its pooled-compression branch. The H3
# checkpoint has no such weight, so the loader zero-initializes it
# (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
@@ -211,11 +257,18 @@ class MiniMaxH3Attention(nn.Module):
query = query.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
key = key.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
value = value.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
if (self.fuse_qknorm_rope and rotary_emb is not None and _can_run_minimax_h3_fusion(query)):
cos, sin = rotary_emb
cos = cos.to(query.dtype)
sin = sin.to(query.dtype)
query = fused_qknorm_rope(query, self.norm_q.weight, cos, sin, self.norm_q.eps)
key = fused_qknorm_rope(key, self.norm_k.weight, cos, sin, self.norm_k.eps)
else:
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
# branch cannot express. Apply it above, then pass no RoPE here.
@@ -397,6 +450,9 @@ class MiniMaxH3TransformerBlock(nn.Module):
quant_config: QuantizationConfig | None,
prefix: str,
adaln_apply_silu: bool = True,
fuse_modulate: bool = False,
fuse_qknorm_rope: bool = False,
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
@@ -408,6 +464,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
supported_attention_backends,
quant_config,
prefix=f"{prefix}.attn",
fuse_qknorm_rope=fuse_qknorm_rope,
)
self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
self.ff = MiniMaxH3FeedForward(
@@ -415,6 +472,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
ffn_dim,
quant_config=quant_config,
prefix=f"{prefix}.ff",
fuse_swiglu=fuse_swiglu,
)
self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
time_embed_dim,
@@ -423,6 +481,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
prefix=f"{prefix}.adaln_proj",
apply_silu=adaln_apply_silu,
)
self.fuse_modulate = fuse_modulate
def forward(
self,
@@ -435,19 +494,39 @@ class MiniMaxH3TransformerBlock(nn.Module):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
t.to(hidden_states.dtype) for t in self.adaln_proj(temb))
residual = hidden_states
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
use_modulate_fusion = self.fuse_modulate and _can_run_minimax_h3_fusion(hidden_states)
if use_modulate_fusion:
norm_hidden_states = fused_rmsnorm_modulate(
hidden_states,
self.norm1.weight,
scale_msa,
shift_msa,
adaln_indices,
self.norm1.eps,
)
else:
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attention_output
residual = hidden_states
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
if use_modulate_fusion:
hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate(
hidden_states,
attention_output,
gate_msa,
self.norm2.weight,
scale_mlp,
shift_mlp,
adaln_indices,
self.norm2.eps,
)
else:
hidden_states = hidden_states + gate_msa.index_select(0, adaln_indices) * attention_output
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
feed_forward_output = self.ff(norm_hidden_states)
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
class MiniMaxH3Transformer3DModel(BaseDiT):
@@ -493,6 +572,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
super().__init__(config, hf_config)
arch = config.arch_config
self.enabled_fusions = _enabled_minimax_h3_fusions()
if self.enabled_fusions:
if HAVE_TRITON:
logger.info(
"MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled and "
"torch.compile-captured forwards fall back to eager).",
",".join(sorted(self.enabled_fusions)))
else:
logger.warning(
"FASTVIDEO_MINIMAX_H3_FUSIONS requested %s but Triton is unavailable; "
"every forward stays on the eager path.", ",".join(sorted(self.enabled_fusions)))
sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1
if arch.num_attention_heads % sp_world_size:
raise ValueError(f"MiniMax H3 attention heads ({arch.num_attention_heads}) must be divisible by "
@@ -546,7 +636,7 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
"parameter, but factorized AdaLN weights are pinned to FP16 "
"(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank "
"checkpoint instead, then re-fit the basis with "
"tools/minimax_h3/fit_adaln_basis.py.")
"scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
adaln_dim = self.adaln_rank or arch.time_embed_dim
self.adaln_basis = ReplicatedLinear(
arch.time_embed_dim,
@@ -590,6 +680,9 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
config.quant_config,
prefix=f"{config.prefix}.transformer_blocks.{index}",
adaln_apply_silu=self.adaln_rank is None,
fuse_modulate="modulate" in self.enabled_fusions,
fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions,
fuse_swiglu="swiglu" in self.enabled_fusions,
) for index in range(arch.num_layers)
])
self.norm_out = MiniMaxH3AdaLayerNormOut(
@@ -616,6 +709,20 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
)
self.__post_init__()
def prepare_for_compile(self) -> None:
"""Pipeline hook, called once right before torch.compile wraps the blocks.
Dynamo capture traces the eager branch of every fusion guard, so an
enabled ``FASTVIDEO_MINIMAX_H3_FUSIONS`` set is silently inert inside
compiled block forwards (H3 compiles per-block by default). Say so
once instead of leaving the flag looking active.
"""
if self.enabled_fusions:
logger.warning(
"torch.compile is enabled for MiniMax H3, so the requested inference fusions (%s) are "
"inert inside compiled block forwards; the compiled eager path runs instead.",
",".join(sorted(self.enabled_fusions)))
def materialize_non_persistent_buffers(
self,
device: torch.device,
@@ -734,14 +841,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
rotary_emb = (rotary_cos, rotary_sin)
for block in self.transformer_blocks:
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
# The eager driver owns profiling markers while each block's compiled
# forward owns the graph that the marker surrounds.
for block_index, block in enumerate(self.transformer_blocks):
with nvtx_range(f"minimax_h3.transformer_block.{block_index}"):
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
packed_hidden_states = self.norm_out(
packed_hidden_states,
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
"""Inference-only MiniMax H3 fusions adapted from NVlabs/Sana Sol-Engine.
Source: https://github.com/NVlabs/Sana/tree/sol-engine/models/minimax_h3/GB200
"""
from .modulation import (
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
)
from .qknorm_rope import HAVE_TRITON, fused_qknorm_rope
from .swiglu import minimax_h3_swiglu
__all__ = [
"HAVE_TRITON",
"fused_qknorm_rope",
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
"minimax_h3_swiglu",
]
@@ -0,0 +1,302 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax H3 RMSNorm and row-indexed modulation fusions."""
from __future__ import annotations
import math
import torch
try:
import triton
import triton.language as tl
except ImportError as exc: # pragma: no cover - depends on the runtime image
triton = None
tl = None
_TRITON_IMPORT_ERROR: ImportError | None = exc
else:
_TRITON_IMPORT_ERROR = None
__all__ = [
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
]
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
_rmsnorm_modulate_kernel = None
_residual_gate_rmsnorm_modulate_kernel = None
if triton is not None:
@triton.jit
def _rmsnorm_modulate_kernel(
out_ptr,
x_ptr,
weight_ptr,
scale_ptr,
shift_ptr,
index_ptr,
n_cols,
n_index,
eps,
stride_x_row,
stride_scale_row,
stride_shift_row,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK)
mask = cols < n_cols
x_offsets = row * stride_x_row + cols
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
x = tl.load(x_ptr + x_offsets, mask=mask, other=0.0).to(tl.float32)
variance = tl.sum(x * x, axis=0) / n_cols
normed = x * tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
scale = tl.load(
scale_ptr + table_row * stride_scale_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
shift = tl.load(
shift_ptr + table_row * stride_shift_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
output = normed * weight * (1.0 + scale) + shift
tl.store(out_ptr + x_offsets, output.to(out_ptr.dtype.element_ty), mask=mask)
@triton.jit
def _residual_gate_rmsnorm_modulate_kernel(
hidden_out_ptr,
normed_out_ptr,
residual_ptr,
branch_ptr,
gate_ptr,
weight_ptr,
scale_ptr,
shift_ptr,
index_ptr,
n_cols,
n_index,
eps,
stride_input_row,
stride_gate_row,
stride_scale_row,
stride_shift_row,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK)
mask = cols < n_cols
input_offsets = row * stride_input_row + cols
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
residual = tl.load(residual_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
branch = tl.load(branch_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(
gate_ptr + table_row * stride_gate_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
hidden = residual + gate * branch
tl.store(hidden_out_ptr + input_offsets, hidden.to(hidden_out_ptr.dtype.element_ty), mask=mask)
variance = tl.sum(hidden * hidden, axis=0) / n_cols
normed = hidden * tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
scale = tl.load(
scale_ptr + table_row * stride_scale_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
shift = tl.load(
shift_ptr + table_row * stride_shift_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
output = normed * weight * (1.0 + scale) + shift
tl.store(normed_out_ptr + input_offsets, output.to(normed_out_ptr.dtype.element_ty), mask=mask)
def _validate_contract(
x: torch.Tensor,
weight: torch.Tensor,
tables: tuple[torch.Tensor, ...],
index: torch.Tensor,
eps: float,
) -> None:
if x.ndim < 2:
raise ValueError(f"x must have shape (..., sequence_length, hidden_size), got {tuple(x.shape)}.")
if x.numel() == 0 or x.shape[-1] == 0:
raise ValueError("x must not be empty.")
if x.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"x must use float16, bfloat16, or float32, got {x.dtype}.")
hidden_size = x.shape[-1]
sequence_length = x.shape[-2]
if weight.shape != (hidden_size, ):
raise ValueError(f"weight must have shape ({hidden_size},), got {tuple(weight.shape)}.")
if weight.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"weight must use float16, bfloat16, or float32, got {weight.dtype}.")
if index.ndim != 1 or index.numel() != sequence_length:
raise ValueError(
f"index must have shape ({sequence_length},) so it can wrap over batch rows, got {tuple(index.shape)}."
)
if index.dtype not in (torch.int32, torch.int64):
raise TypeError(f"index must use int32 or int64, got {index.dtype}.")
if not isinstance(eps, (float, int)) or isinstance(eps, bool) or not math.isfinite(eps) or eps <= 0:
raise ValueError(f"eps must be a positive finite number, got {eps!r}.")
table_rows = tables[0].shape[0] if tables and tables[0].ndim == 2 else None
for name, table in zip(("gate", "scale", "shift")[-len(tables):], tables, strict=True):
if table.ndim != 2 or table.shape[1] != hidden_size:
raise ValueError(f"{name} must have shape (table_rows, {hidden_size}), got {tuple(table.shape)}.")
if table.shape[0] == 0 or table.shape[0] != table_rows:
raise ValueError("all modulation tables must have the same non-zero row count.")
if table.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"{name} must use float16, bfloat16, or float32, got {table.dtype}.")
tensors = (x, weight, *tables, index)
if any(tensor.device != x.device for tensor in tensors[1:]):
raise ValueError("x, weight, modulation tables, and index must be on the same device.")
def _validate_residual_branch(residual: torch.Tensor, branch: torch.Tensor) -> None:
if branch.shape != residual.shape:
raise ValueError(f"branch must match residual shape {tuple(residual.shape)}, got {tuple(branch.shape)}.")
if branch.dtype != residual.dtype:
raise TypeError(f"branch dtype must match residual dtype {residual.dtype}, got {branch.dtype}.")
if branch.device != residual.device:
raise ValueError("branch and residual must be on the same device.")
def _require_triton_cuda(x: torch.Tensor) -> None:
if triton is None:
detail = f": {_TRITON_IMPORT_ERROR}" if _TRITON_IMPORT_ERROR is not None else ""
raise RuntimeError(f"MiniMax H3 modulation fusion requires Triton{detail}.")
if x.device.type != "cuda":
raise RuntimeError(f"MiniMax H3 modulation fusion requires CUDA tensors, got device {x.device}.")
def _require_forward_only(*tensors: torch.Tensor) -> None:
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in tensors):
raise RuntimeError("MiniMax H3 modulation fusion is forward-only and does not support autograd.")
def _next_power_of_two(value: int) -> int:
return 1 << (value - 1).bit_length()
def _num_warps(block_size: int) -> int:
if block_size >= 8192:
return 16
if block_size >= 2048:
return 8
return 4
def _row_addressable(table: torch.Tensor) -> torch.Tensor:
return table if table.stride(-1) == 1 else table.contiguous()
def fused_rmsnorm_modulate(
x: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> torch.Tensor:
"""Run RMSNorm and row-indexed modulation in one strict Triton kernel.
``index`` values must lie in ``[0, table_rows)``. Unlike eager
``index_select``, the kernel does not raise on out-of-range values (a
device-side bounds check would synchronize); callers are safe by
construction (``timestep_indices * 3 + token_tags``, SP pads with 0).
"""
_validate_contract(x, weight, (scale, shift), index, eps)
_require_forward_only(x, weight, scale, shift)
_require_triton_cuda(x)
hidden_size = x.shape[-1]
flat_x = x.reshape(-1, hidden_size).contiguous()
weight = weight.contiguous()
scale = _row_addressable(scale)
shift = _row_addressable(shift)
index = index.contiguous()
output = torch.empty_like(flat_x)
block_size = _next_power_of_two(hidden_size)
_rmsnorm_modulate_kernel[(flat_x.shape[0], )](
output,
flat_x,
weight,
scale,
shift,
index,
hidden_size,
index.numel(),
eps,
flat_x.stride(0),
scale.stride(0),
shift.stride(0),
BLOCK=block_size,
num_warps=_num_warps(block_size),
)
return output.view_as(x)
def fused_residual_gate_rmsnorm_modulate(
residual: torch.Tensor,
branch: torch.Tensor,
gate: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Fuse residual update, row-indexed gate, RMSNorm, and modulation.
``index`` values must lie in ``[0, table_rows)``; see
:func:`fused_rmsnorm_modulate` for why the wrapper does not check them.
"""
_validate_residual_branch(residual, branch)
_validate_contract(residual, weight, (gate, scale, shift), index, eps)
_require_forward_only(residual, branch, gate, weight, scale, shift)
_require_triton_cuda(residual)
hidden_size = residual.shape[-1]
flat_residual = residual.reshape(-1, hidden_size).contiguous()
flat_branch = branch.reshape(-1, hidden_size).contiguous()
weight = weight.contiguous()
gate = _row_addressable(gate)
scale = _row_addressable(scale)
shift = _row_addressable(shift)
index = index.contiguous()
hidden = torch.empty_like(flat_residual)
modulated = torch.empty_like(flat_residual)
block_size = _next_power_of_two(hidden_size)
_residual_gate_rmsnorm_modulate_kernel[(flat_residual.shape[0], )](
hidden,
modulated,
flat_residual,
flat_branch,
gate,
weight,
scale,
shift,
index,
hidden_size,
index.numel(),
eps,
flat_residual.stride(0),
gate.stride(0),
scale.stride(0),
shift.stride(0),
BLOCK=block_size,
num_warps=_num_warps(block_size),
)
return hidden.view_as(residual), modulated.view_as(residual)
@@ -0,0 +1,174 @@
# SPDX-License-Identifier: Apache-2.0
"""Fused per-head RMSNorm and partial rotary embedding for MiniMax H3."""
from __future__ import annotations
import math
import torch
try:
import triton
import triton.language as tl
HAVE_TRITON = True
except ImportError: # pragma: no cover - exercised only in environments without Triton
triton = None
tl = None
HAVE_TRITON = False
if HAVE_TRITON:
@triton.jit
def _qknorm_partial_rope_kernel(
out_ptr,
x_ptr,
weight_ptr,
cos_ptr,
sin_ptr,
head_dim,
rotary_dim,
half_rotary_dim,
num_heads,
seq_len,
eps,
BLOCK_SIZE: tl.constexpr,
):
# int64, like the sibling kernels: with int32 program ids,
# ``row * head_dim`` wraps once the flattened input reaches 2**31
# elements (H3's 56 heads x 128 head_dim crosses that at
# batch*seq >= 299_593 tokens per rank) and the loads/stores below
# become out-of-bounds. ``seq_index`` inherits int64 from ``row``.
row = tl.program_id(0).to(tl.int64)
seq_index = (row // num_heads) % seq_len
cols = tl.arange(0, BLOCK_SIZE)
head_mask = cols < head_dim
row_offset = row * head_dim
x = tl.load(x_ptr + row_offset + cols, mask=head_mask, other=0.0).to(tl.float32)
variance = tl.sum(x * x, axis=0) / head_dim
inv_rms = tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=head_mask, other=0.0).to(tl.float32)
normalized = x * inv_rms * weight
rotary_mask = cols < rotary_dim
first_half = cols < half_rotary_dim
partner_col = tl.where(first_half, cols + half_rotary_dim, cols - half_rotary_dim)
partner_x = tl.load(
x_ptr + row_offset + partner_col,
mask=rotary_mask,
other=0.0,
).to(tl.float32)
partner_weight = tl.load(weight_ptr + partner_col, mask=rotary_mask, other=0.0).to(tl.float32)
partner_normalized = partner_x * inv_rms * partner_weight
rotated = tl.where(first_half, -partner_normalized, partner_normalized)
table_offset = seq_index * rotary_dim + cols
cos = tl.load(cos_ptr + table_offset, mask=rotary_mask, other=1.0).to(tl.float32)
sin = tl.load(sin_ptr + table_offset, mask=rotary_mask, other=0.0).to(tl.float32)
rotary_output = normalized * cos + rotated * sin
output = tl.where(rotary_mask, rotary_output, normalized)
tl.store(out_ptr + row_offset + cols, output.to(out_ptr.dtype.element_ty), mask=head_mask)
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
def _validate_inputs(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> tuple[int, int, int, int, int]:
for name, tensor in (("x", x), ("weight", weight), ("cos", cos), ("sin", sin)):
if not isinstance(tensor, torch.Tensor):
raise TypeError(f"{name} must be a torch.Tensor, got {type(tensor).__name__}")
if x.ndim != 4:
raise ValueError(f"x must have shape (batch, seq, heads, head_dim), got {tuple(x.shape)}")
batch, seq_len, num_heads, head_dim = x.shape
if min(batch, seq_len, num_heads, head_dim) <= 0:
raise ValueError(f"x dimensions must all be positive, got {tuple(x.shape)}")
if weight.shape != (head_dim, ):
raise ValueError(f"weight must have shape ({head_dim},), got {tuple(weight.shape)}")
if cos.ndim != 2:
raise ValueError(f"cos must have shape (seq, rotary_dim), got {tuple(cos.shape)}")
if sin.shape != cos.shape:
raise ValueError(f"sin must match cos shape {tuple(cos.shape)}, got {tuple(sin.shape)}")
if cos.shape[0] != seq_len:
raise ValueError(f"cos/sin sequence length must be {seq_len}, got {cos.shape[0]}")
rotary_dim = cos.shape[1]
if rotary_dim <= 0:
raise ValueError(f"rotary_dim must be positive, got {rotary_dim}")
if rotary_dim > head_dim:
raise ValueError(f"rotary_dim must not exceed head_dim, got rotary_dim={rotary_dim}, head_dim={head_dim}")
if rotary_dim % 2:
raise ValueError(f"rotary_dim must be even, got {rotary_dim}")
if x.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"x dtype must be float16, bfloat16, or float32, got {x.dtype}")
for name, tensor in (("weight", weight), ("cos", cos), ("sin", sin)):
if tensor.dtype != x.dtype:
raise TypeError(f"{name} dtype must match x dtype {x.dtype}, got {tensor.dtype}")
if tensor.device != x.device:
raise ValueError(f"{name} device must match x device {x.device}, got {tensor.device}")
if not isinstance(eps, (float, int)) or not math.isfinite(float(eps)) or eps <= 0:
raise ValueError(f"eps must be a positive finite number, got {eps!r}")
return batch, seq_len, num_heads, head_dim, rotary_dim
def fused_qknorm_rope(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
"""Run per-head RMSNorm and partial RoPE in one Sol-Engine-style kernel.
RMSNorm reduction and RoPE arithmetic stay in FP32 registers until the
final store. Triton's reduction order and the absence of eager's BF16
intermediate materializations can produce small, expected rounding drift.
Row offsets are computed in int64, so inputs beyond 2**31 total elements
(about 300k tokens per rank at H3's 56 heads x 128 head_dim) address
correctly.
"""
batch, seq_len, num_heads, head_dim, rotary_dim = _validate_inputs(x, weight, cos, sin, eps)
if not weight.is_contiguous():
raise ValueError("weight must be contiguous")
if not cos.is_contiguous() or not sin.is_contiguous():
raise ValueError("cos and sin must be contiguous (seq, rotary_dim) tables")
if not x.is_cuda:
raise RuntimeError("fused_qknorm_rope requires CUDA tensors")
if not HAVE_TRITON:
raise RuntimeError("fused_qknorm_rope requires Triton")
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in (x, weight, cos, sin)):
raise RuntimeError("fused_qknorm_rope is inference-only and does not implement autograd")
flat_x = x.reshape(-1, head_dim).contiguous()
flat_out = torch.empty_like(flat_x)
block_size = 1 << (head_dim - 1).bit_length()
_qknorm_partial_rope_kernel[(flat_x.shape[0], )](
flat_out,
flat_x,
weight,
cos,
sin,
head_dim,
rotary_dim,
rotary_dim // 2,
num_heads,
seq_len,
eps,
BLOCK_SIZE=block_size,
num_warps=4,
)
return flat_out.view(batch, seq_len, num_heads, head_dim)
__all__ = ["HAVE_TRITON", "fused_qknorm_rope"]
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax H3's value-first packed SwiGLU fusion."""
from __future__ import annotations
import torch
try:
import triton
import triton.language as tl
HAVE_TRITON = True
except ImportError: # pragma: no cover - exercised only in environments without Triton
triton = None
tl = None
HAVE_TRITON = False
def _validate_input(x: torch.Tensor) -> int:
if x.ndim == 0:
raise ValueError("MiniMax H3 SwiGLU expects at least one dimension")
packed_width = x.shape[-1]
if packed_width == 0 or packed_width % 2 != 0:
raise ValueError(
"MiniMax H3 SwiGLU expects a positive even last dimension containing packed (value, gate) halves, "
f"got {packed_width}"
)
if not x.is_floating_point():
raise TypeError(f"MiniMax H3 SwiGLU expects a floating-point tensor, got {x.dtype}")
return packed_width // 2
if HAVE_TRITON:
@triton.jit
def _minimax_h3_swiglu_kernel(
out_ptr,
x_ptr,
ffn_dim,
stride_in_row,
stride_out_row,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK_SIZE)
mask = cols < ffn_dim
value = tl.load(x_ptr + row * stride_in_row + cols, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(x_ptr + row * stride_in_row + ffn_dim + cols, mask=mask, other=0.0).to(tl.float32)
# Match Sol-Engine: keep the complete SwiGLU expression in FP32 and
# convert only the final output store.
out = value * (gate * tl.sigmoid(gate))
tl.store(out_ptr + row * stride_out_row + cols, out.to(out_ptr.dtype.element_ty), mask=mask)
else:
_minimax_h3_swiglu_kernel = None
def _num_warps(block_size: int) -> int:
if block_size >= 8192:
return 16
if block_size >= 2048:
return 8
return 4
def minimax_h3_swiglu(x: torch.Tensor) -> torch.Tensor:
"""Run the forward-only Triton fusion over an H3 ``(..., 2 * ffn_dim)`` input.
This is intentionally a strict kernel wrapper: callers own fallback policy and
must only invoke it for a supported CUDA inference path.
"""
ffn_dim = _validate_input(x)
if not x.is_cuda:
raise ValueError("MiniMax H3 fused SwiGLU requires a CUDA tensor")
if x.dtype not in (torch.float16, torch.bfloat16, torch.float32):
raise TypeError(f"MiniMax H3 fused SwiGLU supports float16, bfloat16, and float32, got {x.dtype}")
if torch.is_grad_enabled() and x.requires_grad:
raise RuntimeError("MiniMax H3 fused SwiGLU is forward-only and does not implement autograd")
if _minimax_h3_swiglu_kernel is None:
raise RuntimeError("MiniMax H3 fused SwiGLU requires Triton")
packed_width = x.shape[-1]
flat = x.reshape(-1, packed_width).contiguous()
output_shape = (*x.shape[:-1], ffn_dim)
if flat.shape[0] == 0:
return torch.empty(output_shape, dtype=x.dtype, device=x.device)
out = torch.empty((flat.shape[0], ffn_dim), dtype=x.dtype, device=x.device)
block_size = triton.next_power_of_2(ffn_dim)
_minimax_h3_swiglu_kernel[(flat.shape[0],)](
out,
flat,
ffn_dim,
flat.stride(0),
out.stride(0),
BLOCK_SIZE=block_size,
num_warps=_num_warps(block_size),
)
return out.view(output_shape)
__all__ = ["HAVE_TRITON", "minimax_h3_swiglu"]
+8 -8
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from dataclasses import field
from typing import Any, Generic, TypeVar
import torch
from torch import nn
@@ -8,11 +9,16 @@ from torch import nn
from fastvideo.configs.models.encoders import (BaseEncoderOutput, ImageEncoderConfig, TextEncoderConfig)
from fastvideo.platforms import AttentionBackendEnum
TextEncoderOutputT = TypeVar("TextEncoderOutputT")
class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
"""Base for native encoders with a model-specific forward output contract."""
class TextEncoder(nn.Module, ABC):
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = TextEncoderConfig()._supported_attention_backends
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -23,13 +29,7 @@ class TextEncoder(nn.Module, ABC):
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
@abstractmethod
def forward(self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput:
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
pass
@property
@@ -0,0 +1,453 @@
# SPDX-License-Identifier: Apache-2.0
"""Serialized block-FP8 execution for the MiniMax-H3 Qwen3-VL encoder."""
from typing import Any
import torch
from torch import nn
from torch.nn.parameter import Parameter
try:
import triton
import triton.language as tl
except ImportError:
triton = None
tl = None
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.linear import LinearBase, LinearMethodBase
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.layers.quantization.fp8_config import FP8_DTYPE
from fastvideo.models.utils import set_weight_attrs
class MiniMaxH3SerializedFP8Config(QuantizationConfig):
"""Serialized 128x128 block-FP8 contract for the H3 text encoder."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
if weight_block_size != (128, 128):
raise ValueError("MiniMax-H3 serialized FP8 requires weight_block_size=[128, 128], "
f"got {list(weight_block_size)}")
self.weight_block_size = weight_block_size
self.is_checkpoint_fp8_serialized = True
self.activation_scheme = "dynamic"
@classmethod
def get_name(cls) -> str:
return "fp8"
@classmethod
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
return [torch.bfloat16]
@classmethod
def get_min_capability(cls) -> int:
return 100
@staticmethod
def get_config_filenames() -> list[str]:
return []
@classmethod
def from_config(cls, config: dict[str, Any]) -> "MiniMaxH3SerializedFP8Config":
quant_method = str(config.get("quant_method", "")).lower()
if quant_method != "fp8":
raise ValueError(f"MiniMax-H3 only supports serialized FP8 text-encoder checkpoints, got {quant_method!r}")
if str(config.get("activation_scheme", "")).lower() != "dynamic":
raise ValueError("MiniMax-H3 serialized FP8 requires dynamic activation quantization")
if str(config.get("fmt", "e4m3")).lower() not in ("e4m3", "float8_e4m3fn"):
raise ValueError(f"MiniMax-H3 serialized FP8 requires E4M3 weights, got {config.get('fmt')!r}")
block_size = config.get("weight_block_size")
if not isinstance(block_size, list | tuple) or len(block_size) != 2:
raise ValueError("MiniMax-H3 serialized FP8 requires a two-dimensional weight_block_size")
ignored_layers = config.get("modules_to_not_convert", config.get("ignored_layers", []))
if not isinstance(ignored_layers, list | tuple):
raise ValueError("MiniMax-H3 serialized FP8 modules_to_not_convert must be a sequence")
language_exclusions = [
name for name in ignored_layers
if isinstance(name, str) and (name.startswith("language_model.") or ".language_model." in name)
]
if language_exclusions:
raise ValueError("MiniMax-H3 does not support partially quantized language stacks; "
f"ignored language layers: {language_exclusions[:3]}")
if not any(isinstance(name, str) and "visual" in name for name in ignored_layers):
raise ValueError("MiniMax-H3 serialized FP8 requires the vision stack to be listed in "
"modules_to_not_convert")
return cls((int(block_size[0]), int(block_size[1])))
def validate_runtime(self, device: torch.device) -> None:
if device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires a CUDA device; "
f"got {device.type!r}")
capability = torch.cuda.get_device_capability(device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < self.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{self.get_min_capability()} or newer, got sm{capability_number}")
if capability[0] not in (10, 12):
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 currently adapts SGLang's Blackwell "
f"FlashInfer path; got unsupported sm{capability_number}")
_require_sglang_per_token_group_fp8_quantization()
_get_flashinfer_groupwise_fp8_gemm()
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
if isinstance(layer, LinearBase) and ".language_model.layers." in prefix:
return MiniMaxH3SerializedFP8LinearMethod(self.weight_block_size)
return None
# Copyright 2024 SGLang Team
# Licensed under the Apache License, Version 2.0.
# Adapted from SGLang's per-token-group quantization kernels and Blackwell
# FlashInfer dispatch at commit f99c62063c7dcfcd06784b885dc08cb52cf23865:
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/kernels/ops/quantization/fp8_kernel.py
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/srt/layers/quantization/fp8_utils.py
if triton is not None:
@triton.jit
def _h3_per_token_group_quant_fp8_row_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
scale_ptr += group_id
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
@triton.jit
def _h3_per_token_group_quant_fp8_column_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
input_columns,
scale_column_stride,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
groups_per_row = input_columns // group_size
scale_column = group_id % groups_per_row
scale_row = group_id // groups_per_row
scale_ptr += scale_column * scale_column_stride + scale_row
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
else:
_h3_per_token_group_quant_fp8_row_major = None
_h3_per_token_group_quant_fp8_column_major = None
def _require_sglang_per_token_group_fp8_quantization() -> None:
if (triton is None or _h3_per_token_group_quant_fp8_row_major is None
or _h3_per_token_group_quant_fp8_column_major is None):
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires Triton for SGLang-compatible "
"per-token-group activation quantization")
def _sglang_per_token_group_quant_fp8(
input_tensor: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""SGLang-compatible dynamic FP8 quantization for contiguous 2-D activations."""
_require_sglang_per_token_group_fp8_quantization()
if input_tensor.ndim != 2:
raise ValueError(f"per-token-group FP8 quantization expects 2-D input, got {input_tensor.ndim}-D")
if not input_tensor.is_contiguous():
raise ValueError("per-token-group FP8 quantization requires contiguous input")
if input_tensor.shape[-1] % group_size:
raise ValueError(f"activation width {input_tensor.shape[-1]} is not divisible by group_size={group_size}")
quantized = torch.empty_like(input_tensor, dtype=FP8_DTYPE)
rows, columns = input_tensor.shape
groups_per_row = columns // group_size
if column_major_scales:
scales = torch.empty(
(groups_per_row, rows),
device=input_tensor.device,
dtype=torch.float32,
).permute(1, 0)
else:
scales = torch.empty(
(rows, groups_per_row),
device=input_tensor.device,
dtype=torch.float32,
)
if rows:
num_groups = input_tensor.numel() // group_size
block = triton.next_power_of_2(group_size)
num_warps = min(max(block // 256, 1), 8)
if column_major_scales:
_h3_per_token_group_quant_fp8_column_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
columns,
scales.stride(1),
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
else:
_h3_per_token_group_quant_fp8_row_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
return quantized, scales
def _get_flashinfer_groupwise_fp8_gemm():
try:
from flashinfer.gemm import gemm_fp8_nt_groupwise
except (AttributeError, ImportError) as error:
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires "
"flashinfer.gemm.gemm_fp8_nt_groupwise (validated with flashinfer-python==0.6.8). "
"FastVideo will not re-quantize this checkpoint to tensorwise FP8.") from error
return gemm_fp8_nt_groupwise
def _get_flashinfer_groupwise_backend(device: torch.device) -> str:
capability = torch.cuda.get_device_capability(device)
if capability[0] >= 12:
return "cutlass"
if capability[0] == 10:
return "trtllm"
capability_number = capability[0] * 10 + capability[1]
raise RuntimeError(f"FlashInfer groupwise FP8 requires a Blackwell GPU, got sm{capability_number}")
def _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor: torch.Tensor,
weight: torch.Tensor,
block_size: tuple[int, int],
weight_scale: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
input_2d = input_tensor.view(-1, input_tensor.shape[-1])
output_shape = [*input_tensor.shape[:-1], weight.shape[0]]
backend = _get_flashinfer_groupwise_backend(input_tensor.device)
if input_2d.dtype != torch.bfloat16:
raise RuntimeError("MiniMax-H3 FlashInfer groupwise FP8 requires BF16 activations; "
f"got {input_2d.dtype}. The SGLang FP16 Triton GEMM fallback is not enabled for H3.")
if backend == "trtllm" and input_2d.shape[1] < 256:
raise RuntimeError("MiniMax-H3 FlashInfer TRTLLM groupwise FP8 requires K >= 256; "
f"got K={input_2d.shape[1]}. The SGLang Triton GEMM fallback is not enabled for H3.")
gemm_fp8_nt_groupwise = _get_flashinfer_groupwise_fp8_gemm()
block_n, block_k = block_size
q_input, x_scale = _sglang_per_token_group_quant_fp8(
input_2d,
block_k,
column_major_scales=(backend == "trtllm"),
)
if backend == "cutlass":
m, k = input_2d.shape
n = weight.shape[0]
expected_x_scale_shape = (k // block_k, m)
expected_weight_scale_shape = (k // block_k, n // block_n)
if x_scale.shape == (m, k // block_k):
x_scale = x_scale.transpose(-1, -2).contiguous()
if weight_scale.shape == (n // block_n, k // block_k):
weight_scale = weight_scale.transpose(-1, -2).contiguous()
if x_scale.shape != expected_x_scale_shape or weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scale layout mismatch: "
f"x_scale={tuple(x_scale.shape)}, weight_scale={tuple(weight_scale.shape)}, "
f"expected={expected_x_scale_shape}/{expected_weight_scale_shape}")
if x_scale.dtype != torch.float32 or weight_scale.dtype != torch.float32:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scales must be float32")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale.contiguous(),
weight_scale.contiguous(),
out_dtype=input_2d.dtype,
backend="cutlass",
scale_major_mode="MN",
)
else:
expected_x_scale_shape = (input_2d.shape[0], input_2d.shape[1] // block_k)
expected_weight_scale_shape = (weight.shape[0] // block_n, weight.shape[1] // block_k)
if x_scale.shape != expected_x_scale_shape or x_scale.stride(0) != 1:
raise RuntimeError("FlashInfer TRTLLM block-FP8 activation scale layout mismatch: "
f"shape={tuple(x_scale.shape)}, stride={x_scale.stride()}, "
f"expected column-major {expected_x_scale_shape}")
if weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer TRTLLM block-FP8 weight scale layout mismatch: "
f"shape={tuple(weight_scale.shape)}, expected={expected_weight_scale_shape}")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale,
weight_scale,
out_dtype=input_2d.dtype,
backend="trtllm",
)
if bias is not None:
output += bias
return output.to(dtype=input_2d.dtype).view(*output_shape)
class MiniMaxH3SerializedFP8LinearMethod(LinearMethodBase):
"""Execute serialized 128x128 block-FP8 weights without re-quantizing them."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
self.weight_block_size = weight_block_size
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: list[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
output_size_per_partition = sum(output_partition_sizes)
block_n, block_k = self.weight_block_size
tp_size = get_tp_world_size()
if tp_size > 1 and input_size // input_size_per_partition == tp_size:
if input_size_per_partition % block_k:
raise ValueError(f"Weight input_size_per_partition={input_size_per_partition} is not divisible "
f"by block_k={block_k}")
if tp_size > 1 and output_size // output_size_per_partition == tp_size:
for output_partition_size in output_partition_sizes:
if output_partition_size % block_n:
raise ValueError(f"Weight output_partition_size={output_partition_size} is not divisible "
f"by block_n={block_n}")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
layer.orig_dtype = params_dtype
weight_loader = extra_weight_attrs.get("weight_loader")
weight = Parameter(
torch.empty(output_size_per_partition, input_size_per_partition, dtype=FP8_DTYPE),
requires_grad=False,
)
set_weight_attrs(weight, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
layer.register_parameter("weight", weight)
scale = Parameter(
torch.empty((output_size_per_partition + block_n - 1) // block_n,
(input_size_per_partition + block_k - 1) // block_k,
dtype=torch.float32),
requires_grad=False,
)
set_weight_attrs(scale, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
scale.data.fill_(torch.finfo(torch.float32).min)
layer.register_parameter("weight_scale_inv", scale)
layer.register_parameter("input_scale", None)
def process_weights_after_loading(self, layer: nn.Module) -> None:
weight = getattr(layer, "weight", None)
block_scales = getattr(layer, "weight_scale_inv", None)
if weight is None or block_scales is None:
raise ValueError("Serialized MiniMax-H3 FP8 linear is missing weight or weight_scale_inv")
if weight.dtype != FP8_DTYPE:
raise ValueError(f"Serialized MiniMax-H3 FP8 weight must be {FP8_DTYPE}, got {weight.dtype}")
if block_scales.dtype != torch.float32:
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must be float32, "
f"got {block_scales.dtype}")
block_n, block_k = self.weight_block_size
output_size, input_size = weight.shape
if output_size % block_n or input_size % block_k:
raise ValueError("Serialized MiniMax-H3 FP8 weight dimensions must be divisible by the 128x128 block size; "
f"got {tuple(weight.shape)}")
expected_scale_shape = (output_size // block_n, input_size // block_k)
if tuple(block_scales.shape) != expected_scale_shape:
raise ValueError("Serialized MiniMax-H3 FP8 scale shape mismatch: "
f"expected {expected_scale_shape}, got {tuple(block_scales.shape)}")
if not bool(torch.isfinite(block_scales).all()) or bool((block_scales <= 0).any()):
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must contain finite positive values")
layer.weight.data = weight.data
layer.weight_scale_inv.data = block_scales.data
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
if x.device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 execution requires CUDA")
capability = torch.cuda.get_device_capability(x.device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < MiniMaxH3SerializedFP8Config.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{MiniMaxH3SerializedFP8Config.get_min_capability()} or newer, "
f"got sm{capability_number}")
if not x.is_contiguous():
x = x.contiguous()
return _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
x,
layer.weight,
self.weight_block_size,
layer.weight_scale_inv,
bias,
)
__all__ = [
"MiniMaxH3SerializedFP8Config",
"MiniMaxH3SerializedFP8LinearMethod",
]
@@ -8,13 +8,13 @@ import torch
import torch.nn.functional as F
from torch import nn
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import MiniMaxH3SerializedFP8Config
from fastvideo.models.loader.weight_utils import default_weight_loader
@@ -227,19 +227,13 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
)
# Build only as far as the consumer reads. The hidden-state tuple records
# each layer's input, so stopping after N layers still yields entry N,
# the output of layer N-1, unchanged. Everything above it exists only to
# feed `last_hidden_state`, which nothing consumes.
override = config.num_hidden_layers_override
self.num_layers = (config.num_hidden_layers
if override is None else min(config.num_hidden_layers, override))
self.output_hidden_state_index = config.output_hidden_state_index
self.layers = nn.ModuleList(
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
for index in range(self.num_layers))
# The final norm sits above the tapped layer, so a truncated stack drops
# it. Keeping it would overwrite the tapped entry with a normalised
# tensor and change conditioning without raising anything.
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
if self.num_layers == config.num_hidden_layers else None)
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
@@ -249,18 +243,14 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
inputs_embeds: torch.Tensor,
position_ids: torch.Tensor,
attention_mask: torch.Tensor | None,
output_hidden_states: bool,
visual_pos_masks: torch.Tensor | None,
deepstack_visual_embeds: list[torch.Tensor] | None,
) -> BaseEncoderOutput:
) -> torch.Tensor:
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
attention_mask = None
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
hidden_states = inputs_embeds
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
for layer_index, layer in enumerate(self.layers):
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
hidden_states = layer(hidden_states, position_embeddings, attention_mask)
if deepstack_visual_embeds is not None and layer_index < len(deepstack_visual_embeds):
if visual_pos_masks is None:
@@ -269,13 +259,9 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
visual = deepstack_visual_embeds[layer_index].to(hidden_states.device, hidden_states.dtype)
updated = hidden_states[mask].clone() + visual
hidden_states[mask] = updated
if self.norm is not None:
hidden_states = self.norm(hidden_states)
# Truncated or not, the last entry is appended here, so the tapped index
# lands in the same place either way.
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
if layer_index + 1 == self.output_hidden_state_index:
return hidden_states
raise RuntimeError(f"MiniMax-H3 text stack did not reach hidden_states[{self.output_hidden_state_index}]")
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
@@ -513,10 +499,18 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
return self.merger(hidden_states), deepstack_features
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
supports_hf_from_pretrained = False
supported_checkpoint_quantization_methods = frozenset({"fp8"})
@classmethod
def checkpoint_quantization_config_from_metadata(
cls,
metadata: dict[str, Any],
) -> MiniMaxH3SerializedFP8Config:
return MiniMaxH3SerializedFP8Config.from_config(metadata)
def __init__(self, config: MiniMaxH3Qwen3VLConfig) -> None:
super().__init__(config)
@@ -530,15 +524,12 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
@property
def num_hidden_layers(self) -> int:
"""The checkpoint architecture's nominal depth, matching its config.json.
When ``num_hidden_layers_override`` truncates the stack at the
conditioning tap, fewer layers exist; the built count is
``self.language_model.num_layers``, and the hidden-state tuple has
``num_layers + 1`` entries, not ``num_hidden_layers + 1``.
"""
return self.config.num_hidden_layers
@property
def num_built_hidden_layers(self) -> int:
return self.language_model.num_layers
def _get_rope_index(
self,
input_ids: torch.Tensor,
@@ -631,35 +622,39 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
f"tokens={int(mask.sum())}, features={features.shape[0]}")
return mask
def forward(
# no_grad, NOT inference_mode: with text_encoder_cpu_offload=True (the
# FastVideoArgs default) the loader FSDP2-shards this conditioner, and
# FSDP2's wait_for_unshard reads tensor._version via
# _unsafe_preserve_version_counter - inference tensors do not track
# version counters, so inference_mode crashes the first encode. no_grad
# frees the same activation memory and keeps prompt_embeds ordinary
# tensors (safe for any future backward through the conditioning).
@torch.no_grad()
def encode_ids(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
mm_token_type_ids: torch.Tensor | None = None,
**kwargs: Any,
) -> BaseEncoderOutput:
del mm_token_type_ids, kwargs
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Exactly one of input_ids or inputs_embeds is required")
if inputs_embeds is None:
assert input_ids is not None
inputs_embeds = self.language_model.embed_tokens(input_ids)
if input_ids is None and (pixel_values is not None or pixel_values_videos is not None):
raise ValueError("Multimodal Qwen3-VL inputs require input_ids for placeholder matching")
) -> torch.Tensor:
if input_ids.ndim != 1:
raise ValueError(f"MiniMax-H3 slim forward expects 1-D input_ids, got shape={tuple(input_ids.shape)}")
if (pixel_values is None) != (image_grid_thw is None):
raise ValueError("pixel_values and image_grid_thw must be provided together")
if (pixel_values_videos is None) != (video_grid_thw is None):
raise ValueError("pixel_values_videos and video_grid_thw must be provided together")
input_ids = input_ids.unsqueeze(0)
inputs_embeds = self.language_model.embed_tokens(input_ids)
image_mask = None
video_mask = None
image_deepstack = None
video_deepstack = None
if pixel_values is not None:
if input_ids is None or image_grid_thw is None:
if image_grid_thw is None:
raise ValueError("pixel_values require input_ids and image_grid_thw")
image_features, image_deepstack = self._visual_features(pixel_values, image_grid_thw)
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -667,7 +662,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"image")
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
if pixel_values_videos is not None:
if input_ids is None or video_grid_thw is None:
if video_grid_thw is None:
raise ValueError("pixel_values_videos require input_ids and video_grid_thw")
video_features, video_deepstack = self._visual_features(pixel_values_videos, video_grid_thw)
video_features = video_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -695,50 +690,34 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
visual_mask = video_mask
deepstack_features = video_deepstack
if position_ids is None:
if input_ids is None:
sequence_length = inputs_embeds.shape[1]
position_ids = torch.arange(sequence_length,
device=inputs_embeds.device).view(1, 1,
-1).expand(3, inputs_embeds.shape[0], -1)
else:
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, attention_mask)
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
outputs = self.language_model(
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
hidden_states = self.language_model(
inputs_embeds,
position_ids,
attention_mask,
output_hidden_states,
None,
visual_mask,
deepstack_features,
)
outputs.attention_mask = attention_mask
return outputs
if hidden_states.ndim != 3 or hidden_states.shape[0] != 1:
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
return hidden_states[0]
def _is_above_the_tap(self, name: str) -> bool:
"""Whether this checkpoint key belongs to a layer we did not build.
A truncated language stack still ships every layer in the checkpoint, and
the unexpected-key check below is strict on purpose, so the surplus keys
have to be dropped here rather than by relaxing it.
"""
language_model = self.language_model
# The final norm is dropped exactly when the stack is truncated, so its
# absence is the signal.
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
if not index.isdigit():
return False
# Only drop indexes the full stack would have built. Anything at or
# above the checkpoint's own num_hidden_layers is corrupt and must
# still raise below, exactly as it does without truncation.
return language_model.num_layers <= int(index) < self.config.num_hidden_layers
def forward(
self,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
return self.encode_ids(
input_ids,
pixel_values=pixel_values,
image_grid_thw=image_grid_thw,
pixel_values_videos=pixel_values_videos,
video_grid_thw=video_grid_thw,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
parameters = dict(self.named_parameters())
@@ -748,7 +727,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
if source_name == "lm_head.weight":
continue
name = source_name[6:] if source_name.startswith("model.") else source_name
if self._is_above_the_tap(name):
if self._is_omitted_checkpoint_key(name):
continue
if name not in parameters:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
@@ -758,7 +737,23 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
loaded.add(name)
return loaded
def _is_omitted_checkpoint_key(self, name: str) -> bool:
"""Return whether a valid checkpoint key belongs to an unbuilt layer."""
language_model = self.language_model
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
return (index.isdigit() and language_model.num_layers <= int(index) < self.config.num_hidden_layers)
EntryClass = MiniMaxH3Qwen3VLConditioner
__all__ = ["MiniMaxH3Qwen3VLConditioner"]
__all__ = [
"MiniMaxH3Qwen3VLConditioner",
"MiniMaxH3SerializedFP8Config",
]
+61 -15
View File
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from contextlib import nullcontext
from copy import deepcopy
from typing import cast
from typing import Any, cast
import torch
import torch.distributed as dist
@@ -30,9 +30,13 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.quantization import get_quantization_config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_resolve_text_encoder_checkpoint_path,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files,
@@ -347,22 +351,46 @@ class TextEncoderLoader(ComponentLoader):
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
runtime_device = get_local_torch_device()
from fastvideo.platforms import current_platform
if cpu_offload:
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
# Set quantization config if specified
if (use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None):
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
checkpoint_path = _resolve_text_encoder_checkpoint_path(
model_path,
fastvideo_args,
use_text_encoder_override,
)
checkpoint_quant_config = _configure_text_encoder_quantization(
model_config,
model_cls,
checkpoint_path,
)
if checkpoint_quant_config is not None:
if fastvideo_args.override_text_encoder_quant is not None:
raise ValueError("Serialized checkpoint quantization is selected from checkpoint metadata; "
"override_text_encoder_quant is an online conversion option and must be unset")
requested_dtype = PRECISION_TO_TYPE[dtype]
if requested_dtype not in checkpoint_quant_config.get_supported_act_dtypes():
raise ValueError(f"Serialized {checkpoint_quant_config.get_name()} text encoder does not support "
f"activation dtype {requested_dtype}")
checkpoint_quant_config.validate_runtime(runtime_device)
logger.info(
"Selected serialized %s text-encoder checkpoint execution from %s",
checkpoint_quant_config.get_name(),
checkpoint_path,
)
elif use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
@@ -381,11 +409,20 @@ class TextEncoderLoader(ComponentLoader):
weights_to_load = {name for name, _ in model.named_parameters()}
if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None):
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
if os.path.isdir(checkpoint_path):
override_weights = self._get_all_weights(
model,
checkpoint_path,
to_cpu=bool(cpu_offload),
)
else:
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
override_weights = safetensors_weights_iterator(
[checkpoint_path],
to_cpu=use_cpu_offload,
)) # type: ignore
)
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
@@ -400,6 +437,10 @@ class TextEncoderLoader(ComponentLoader):
self.counter_after_loading_weights - self.counter_before_loading_weights,
)
if checkpoint_quant_config is not None:
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
logger.info("Validated %d serialized blockwise FP8 text-encoder linears", processed_linears)
# Explicitly move model to target device after loading weights
model = model.to(target_device)
@@ -442,7 +483,7 @@ class TextEncoderLoader(ComponentLoader):
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded and model_config.quant_config is None:
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
@@ -1057,7 +1098,12 @@ class TransformerLoader(ComponentLoader):
# so recording here makes the decision readable from the loaded
# transformer — and records the narrowed one for teacher/critic.
resolved = record_resolved_attention_backend(dit_config)
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
# Every worker records its resolved backend so distributed profile
# snapshots can prove that all ranks use the requested kernels.
logger.info("Worker %s transformer attention backend: %s",
os.environ.get("RANK", "0"),
resolved.name if resolved else "automatic selection",
local_main_process_only=False)
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
"""Checkpoint-serialized quantization lifecycle for native text encoders."""
import json
import os
from itertools import chain
from typing import Any
import torch
import torch.nn as nn
from safetensors.torch import safe_open
from fastvideo.configs.models import EncoderConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.linear import LinearBase, UnquantizedLinearMethod
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.models.encoders.base import TextEncoder
def _resolve_text_encoder_checkpoint_path(
model_path: str,
fastvideo_args: FastVideoArgs,
use_text_encoder_override: bool,
) -> str:
override = fastvideo_args.override_text_encoder_safetensors if use_text_encoder_override else None
checkpoint_path = override or model_path
if not os.path.exists(checkpoint_path):
raise FileNotFoundError(f"Text-encoder checkpoint does not exist: {checkpoint_path}")
if not os.path.isdir(checkpoint_path) and not os.path.isfile(checkpoint_path):
raise ValueError(f"Text-encoder checkpoint must be a file or directory: {checkpoint_path}")
return checkpoint_path
def _read_text_encoder_checkpoint_quantization_config(checkpoint_path: str) -> dict[str, Any] | None:
checkpoint_dir = checkpoint_path if os.path.isdir(checkpoint_path) else os.path.dirname(checkpoint_path)
config_path = os.path.join(checkpoint_dir, "config.json")
if os.path.isfile(config_path):
try:
with open(config_path, encoding="utf-8") as config_file:
checkpoint_config = json.load(config_file)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid text-encoder checkpoint config: {config_path}") from error
quantization_config = checkpoint_config.get("quantization_config")
if quantization_config is not None:
if not isinstance(quantization_config, dict):
raise ValueError(f"quantization_config in {config_path} must be an object")
return quantization_config
if not os.path.isfile(checkpoint_path) or not checkpoint_path.endswith(".safetensors"):
return None
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint_file:
metadata = checkpoint_file.metadata() or {}
for key in ("quantization_config", "_quantization_metadata"):
serialized = metadata.get(key)
if serialized is None:
continue
try:
quantization_config = json.loads(serialized)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid {key} metadata in {checkpoint_path}") from error
if not isinstance(quantization_config, dict):
raise ValueError(f"{key} metadata in {checkpoint_path} must decode to an object")
return quantization_config
return None
def _configure_text_encoder_quantization(
model_config: EncoderConfig,
model_cls: type[nn.Module],
checkpoint_path: str,
) -> QuantizationConfig | None:
if not issubclass(model_cls, TextEncoder):
return None
checkpoint_quantization = _read_text_encoder_checkpoint_quantization_config(checkpoint_path)
if checkpoint_quantization is None:
return None
quant_method = str(checkpoint_quantization.get("quant_method", "")).lower()
if not quant_method:
raise ValueError(f"Quantized text-encoder checkpoint {checkpoint_path} does not declare quant_method")
supported_methods = getattr(model_cls, "supported_checkpoint_quantization_methods", frozenset())
if quant_method not in supported_methods:
supported = ", ".join(sorted(supported_methods)) or "none"
raise ValueError(f"Text encoder {model_cls.__name__} does not support serialized {quant_method!r} "
f"checkpoints (supported: {supported})")
factory = getattr(model_cls, "checkpoint_quantization_config_from_metadata", None)
if not callable(factory):
raise ValueError(f"Text encoder {model_cls.__name__} advertises serialized {quant_method!r} support "
"without a checkpoint quantization factory")
quant_config = factory(checkpoint_quantization)
model_config.quant_config = quant_config
return quant_config
def _module_tensor_device(module: nn.Module) -> torch.device | None:
devices = {
tensor.device
for tensor in chain(
module.parameters(recurse=False),
module.buffers(recurse=False),
)
}
if len(devices) > 1:
raise ValueError(f"Quantized text-encoder module {type(module).__name__} spans multiple devices: {devices}")
return next(iter(devices), None)
def _process_quantized_text_encoder_weights(model: nn.Module, process_device: torch.device) -> int:
"""Run quantized post-load hooks one linear at a time on ``process_device``."""
processed = 0
for module in model.modules():
if not isinstance(module, LinearBase) or isinstance(module.quant_method, UnquantizedLinearMethod):
continue
if module.quant_method is None:
continue
original_device = _module_tensor_device(module)
try:
module.to(process_device)
module.quant_method.process_weights_after_loading(module)
finally:
if original_device is not None:
module.to(original_device)
processed += 1
if processed == 0:
raise ValueError("Serialized quantized text-encoder checkpoint selected, but no quantized linear layers exist")
return processed
@@ -392,9 +392,16 @@ class MiniMaxH3AudioBigVGANDecoder(nn.Module):
return torch.clamp(hidden_states, min=-1.0, max=1.0)
def _is_minimax_h3_audio_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the audio decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3AudioBigVGANDecoder)
class MiniMaxH3AudioVAE(nn.Module):
"""DAC encoder plus BigVGAN decoder for mono 32 kHz waveforms."""
_compile_conditions = [_is_minimax_h3_audio_vae_decoder]
def __init__(self, config: MiniMaxH3AudioVAEConfig):
super().__init__()
self.config = config
@@ -0,0 +1,378 @@
# SPDX-License-Identifier: Apache-2.0
"""Sequence-parallel chunk scheduling for the MiniMax-H3 video VAE.
The H3 video VAE decodes a video as a series of temporal-chunk decoder
forwards whose outputs are joined by a short deterministic frame blend
(``AutoencoderKLMiniMaxH3._decode_chunks``), and encodes videos as fully
independent ``clip_length``-frame encoder forwards. Neither the chunk decode
nor the clip encode has any cross-chunk data dependency — only the *joining*
of decoded chunks (overlap blending, frame trimming) is sequential. This
module round-robins the chunk/clip forwards across the ranks of a
sequence-parallel group and replays the serial joining logic on the
assembling rank, reproducing the serial result bit for bit.
Bit-exactness contract:
- every rank holds an identical copy of the inputs (the H3 DiT all-gathers
its outputs, and reference pixels are prepared identically on all ranks);
- a chunk decoded on any rank is bitwise the tensor the serial loop would
produce (identical weights, inputs, and deterministic kernels on identical
GPUs), and NCCL transports it bitwise;
- every serialization point of the serial algorithm (overlap blending, frame
trimming, pixel denormalization, output-buffer copies, moment
concatenation and token-drop trimming) runs on the assembling rank in
serial order via the same VAE methods the serial path uses.
Collective safety: all group ranks must call these functions together with
identically shaped inputs. Work proceeds in rounds of one collective each;
ranks without a chunk in the final round contribute a placeholder tensor, so
participation is uniform by construction and no rank-dependent branch guards
a collective.
Caveat — compiled decoders (``enable_torch_compile_vae``): inductor autotunes
kernel configs per process at first call, so a compiled decoder is only
deterministic WITHIN a process, not across processes. Chunks decoded on other
ranks then differ from the serial rank's decode of the same chunk exactly as
two serial runs in different processes would (measured on GB200 at 124f:
max 63/255 on <0.5% of pixels, mean ~1e-2/255, first chunk bit-identical).
With the eager decoder — the pipeline default — parallel output is bitwise
equal to serial ``decode_to_pixels``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from fastvideo.models.vaes.minimax_h3_video import (
AutoencoderKLMiniMaxH3,
AutoencoderKLOutput,
DiagonalGaussianDistribution,
)
from fastvideo.profiler import nvtx_range
if TYPE_CHECKING:
from fastvideo.distributed.parallel_state import GroupCoordinator
# Collective used to move decoded chunk segments to the assembling rank.
# "gather" moves each segment once (destination-only); "all_gather" also
# leaves every rank with every segment. Both are exact; the default is the
# faster one measured on GB200 NVL72 (see the PR notes).
DECODE_GATHER_STRATEGIES = ("gather", "all_gather")
DEFAULT_DECODE_GATHER_STRATEGY = "gather"
def parallel_chunk_indices(num_chunks: int, world_size: int, rank_in_group: int) -> list[int]:
"""Round-robin chunk ownership: chunk ``i`` belongs to rank ``i % world_size``."""
if num_chunks < 0:
raise ValueError(f"num_chunks must be non-negative, got {num_chunks}.")
if world_size < 1:
raise ValueError(f"world_size must be positive, got {world_size}.")
if not 0 <= rank_in_group < world_size:
raise ValueError(f"rank_in_group {rank_in_group} out of range for world_size {world_size}.")
return list(range(rank_in_group, num_chunks, world_size))
def _num_rounds(num_chunks: int, world_size: int) -> int:
return -(-num_chunks // world_size)
def _decode_segment(vae: AutoencoderKLMiniMaxH3, z_padded: torch.Tensor, chunk_index: int) -> torch.Tensor:
"""Decode one temporal chunk's clip and keep the frames the join consumes.
The serial loop uses two spans of each decoded clip: the chunk body
``clip[:, :, frame_pre_padding:chunk_num_frames]`` and (when
``token_drop > 0``) the blend tail
``clip[:, :, chunk_num_frames + frame_pre_padding:]``. Everything from
``frame_pre_padding`` on covers both, so one contiguous slice per chunk
travels over the wire. ``.contiguous()`` also detaches the segment from
any decoder-owned storage (e.g. a compiled decoder's reuse pools) before
the next chunk decode can overwrite it.
"""
start = chunk_index * vae.tokens_chunk_size
with nvtx_range(f"minimax_h3.vae.parallel_chunk.{chunk_index}"):
clip = vae._decode_clip(z_padded[:, :, start:start + vae.tokens_chunk_size + vae.token_overlap])
return clip[:, :, vae.frame_pre_padding:].contiguous()
class _ChunkAssembler:
"""Replay the serial chunk-joining semantics of ``_decode_chunks`` +
``_decode_to_pixels`` on gathered chunk segments, in chunk order.
On CUDA the joining kernels and output copies run on a dedicated side
stream: they depend only on already-gathered segments, so running them
off the main stream keeps the assembling rank's next chunk decode (and
therefore every other rank's next collective) off the assembly's tail.
Stream placement cannot change values — the ops and their order are
identical — so bit-exactness with the serial path is unaffected.
"""
def __init__(self, vae: AutoencoderKLMiniMaxH3, output: torch.Tensor, output_num_frames: int,
non_blocking: bool, device: torch.device) -> None:
self._vae = vae
self._output = output
self._output_num_frames = output_num_frames
self._non_blocking = non_blocking
self._body_frames = vae.tokens_chunk_size * vae.temporal_compression_ratio - vae.frame_pre_padding
self._overlap: torch.Tensor | None = None
self._frame_start = 0
self._stream = torch.cuda.Stream(device) if device.type == "cuda" else None
def push(self, segment: torch.Tensor) -> None:
"""Consume the next chunk's segment (``clip[:, :, frame_pre_padding:]``)."""
if self._stream is None:
self._push(segment)
return
# The segment is produced on the current (collective) stream; hand it
# to the assembly stream and pin its storage until assembly reads it.
self._stream.wait_stream(torch.cuda.current_stream(segment.device))
segment.record_stream(self._stream)
with torch.cuda.stream(self._stream):
self._push(segment)
def _push(self, segment: torch.Tensor) -> None:
vae = self._vae
chunk = segment[:, :, :self._body_frames]
if self._overlap is not None:
chunk = vae._blend(self._overlap, chunk, vae.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], self._output_num_frames - self._frame_start)
chunk = chunk[:, :, :num_frames]
# The tail past the body (and its pre-padding gap) is the next
# chunk's blend overlap — the serial loop's ``next_overlap``.
self._overlap = segment[:, :, self._body_frames + vae.frame_pre_padding:] if vae.config.token_drop > 0 else None
if num_frames > 0:
self._emit(chunk)
def finalize(self) -> None:
"""Emit the final overlap tail exactly as the serial generator does."""
if self._overlap is not None and self._frame_start < self._output_num_frames:
tail = self._overlap[:, :, :self._output_num_frames - self._frame_start]
if self._stream is None:
self._emit(tail)
else:
with torch.cuda.stream(self._stream):
self._emit(tail)
if self._frame_start != self._output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {self._frame_start} frames into an output buffer expecting "
f"{self._output.shape[2]}.")
def synchronize(self) -> None:
"""Drain assembly kernels and output copies before the buffer is read."""
if self._stream is not None:
self._stream.synchronize()
def _emit(self, chunk: torch.Tensor) -> None:
pixels = self._vae.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._vae._copy_chunk_pixels(pixels, self._output, self._frame_start, self._non_blocking)
self._frame_start += pixels.shape[2]
def _broadcast_segment_meta(group: "GroupCoordinator",
segment: torch.Tensor | None) -> tuple[torch.dtype, tuple[int, ...]]:
"""Share the leader's real segment dtype/shape so placeholder tensors match.
The decoder's output dtype depends on the surrounding autocast context;
deriving it on the leader from an actually decoded segment (instead of
predicting it) keeps collective dtypes correct by construction.
"""
meta = (segment.dtype, tuple(segment.shape)) if segment is not None else None
meta = group.broadcast_object(meta, src=0)
if meta is None:
raise RuntimeError("MiniMax-H3 parallel VAE meta broadcast returned no leader metadata.")
return meta
def decode_to_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str = DEFAULT_DECODE_GATHER_STRATEGY,
) -> torch.Tensor | None:
"""Chunk-parallel ``decode_to_pixels`` across a sequence-parallel group.
All group ranks call this together with identical ``z``. Temporal chunks
are decoded round-robin across the group and their segments move to the
group's first rank, which assembles bitwise the serial
``decode_to_pixels`` result into ``output``. Only the first rank passes
``output`` (validated exactly like the serial API); other ranks pass
``None`` and receive ``None``.
"""
if strategy not in DECODE_GATHER_STRATEGIES:
raise ValueError(f"Unknown parallel-decode strategy {strategy!r}; expected one of {DECODE_GATHER_STRATEGIES}.")
is_leader = group.rank_in_group == 0
if is_leader:
if output is None:
raise ValueError("The first sequence-parallel rank must provide the CPU output buffer.")
expected_shape = vae.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
elif output is not None:
raise ValueError("Only the first sequence-parallel rank may provide an output buffer.")
if group.world_size == 1:
return vae.decode_to_pixels(z, output)
try:
if vae.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
slice_output = output[batch_index:batch_index + 1] if output is not None else None
_decode_single_parallel(vae, z_slice, slice_output, group, strategy)
else:
_decode_single_parallel(vae, z, output, group, strategy)
finally:
# Drain the leader's async chunk copies before the caller (or an
# exception handler) can read or release the pinned buffer.
if output is not None and vae._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def _decode_single_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str,
) -> None:
pad_tokens, num_chunks, output_num_frames = vae._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
world_size = group.world_size
rank = group.rank_in_group
# Every rank decodes its round-0 chunk BEFORE the metadata rendezvous so
# the first decodes run concurrently (a rank that waited on the broadcast
# first would idle a full chunk-decode behind the leader). The leader
# owns chunk 0 under round-robin assignment, so its segment supplies real
# dtype/shape for placeholder rounds instead of guessing autocast state.
first_segment = _decode_segment(vae, z, rank) if rank < num_chunks else None
segment_dtype, segment_shape = _broadcast_segment_meta(group, first_segment if rank == 0 else None)
assembler = None
if output is not None:
non_blocking = vae._streams_chunk_copies(z, output)
assembler = _ChunkAssembler(vae, output, output_num_frames, non_blocking, z.device)
try:
segment_frames = segment_shape[2]
for round_index in range(_num_rounds(num_chunks, world_size)):
chunk_index = round_index * world_size + rank
if chunk_index >= num_chunks:
segment = torch.zeros(segment_shape, dtype=segment_dtype, device=z.device)
elif round_index == 0 and first_segment is not None:
segment = first_segment
else:
segment = _decode_segment(vae, z, chunk_index)
with nvtx_range(f"minimax_h3.vae.parallel_{strategy}.{round_index}"):
if strategy == "gather":
gathered = group.gather(segment, dst=0, dim=2)
else:
gathered = group.all_gather(segment, dim=2)
if assembler is None or gathered is None:
continue
for slot in range(world_size):
if round_index * world_size + slot >= num_chunks:
break
assembler.push(gathered.narrow(2, slot * segment_frames, segment_frames))
if assembler is not None:
assembler.finalize()
finally:
# Drain assembly-stream copies into ``output`` even on the error path
# so an exception cannot leave an in-flight DMA into a buffer the
# caller may release.
if assembler is not None:
assembler.synchronize()
def _encode_clip_moments(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor, clip_index: int) -> torch.Tensor:
"""Encode one ``clip_length``-frame clip exactly as ``_encode_pixels`` does."""
clip_length = vae.config.clip_length
frame_start = clip_index * clip_length
with nvtx_range(f"minimax_h3.vae.parallel_encode_clip.{clip_index}"):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=vae.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = vae.normalize_pixels(clip)
return vae._encode_clip(clip).contiguous()
def encode_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
pixels: torch.Tensor,
group: "GroupCoordinator",
) -> AutoencoderKLOutput:
"""Clip-parallel ``encode_pixels`` across a sequence-parallel group.
Encoder clips have no cross-clip dependency (no overlap, no blending), so
ranks encode disjoint clips and all-gather the per-clip moment tensors.
Every rank returns the identical full posterior — preserving the serial
contract that all ranks hold the same encoded latents — bitwise equal to
``vae.encode_pixels(pixels)``. Moments are latent-sized (a few MB per
clip), so the all-gather is negligible next to the clip forwards.
"""
if pixels.ndim != 5 or pixels.shape[1] != vae.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {vae.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if group.world_size == 1:
return vae.encode_pixels(pixels)
if vae.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([_encode_single_parallel(vae, pixel_slice, group) for pixel_slice in pixels.split(1)])
else:
moments = _encode_single_parallel(vae, pixels, group)
return AutoencoderKLOutput(latent_dist=DiagonalGaussianDistribution(moments))
def _encode_single_parallel(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor,
group: "GroupCoordinator") -> torch.Tensor:
clip_length = vae.config.clip_length
num_clips = -(-pixels.shape[2] // clip_length)
world_size = group.world_size
rank = group.rank_in_group
# Same first-work-then-rendezvous ordering as the decode path: encode the
# round-0 clip before the metadata broadcast so first encodes overlap.
first_moments = _encode_clip_moments(vae, pixels, rank) if rank < num_clips else None
moment_dtype, moment_shape = _broadcast_segment_meta(group, first_moments if rank == 0 else None)
moment_tokens = moment_shape[2]
parts: list[torch.Tensor] = []
for round_index in range(_num_rounds(num_clips, world_size)):
clip_index = round_index * world_size + rank
if clip_index >= num_clips:
moments = torch.zeros(moment_shape, dtype=moment_dtype, device=vae.pixel_mean.device)
elif round_index == 0 and first_moments is not None:
moments = first_moments
else:
moments = _encode_clip_moments(vae, pixels, clip_index)
gathered = group.all_gather(moments, dim=2)
for slot in range(world_size):
if round_index * world_size + slot >= num_clips:
break
parts.append(gathered.narrow(2, slot * moment_tokens, moment_tokens))
encoded = torch.cat(parts, dim=2)
if vae.config.token_drop > 0:
encoded = encoded[:, :, :-vae.config.token_drop]
return encoded
__all__ = [
"DECODE_GATHER_STRATEGIES",
"DEFAULT_DECODE_GATHER_STRATEGY",
"decode_to_pixels_parallel",
"encode_pixels_parallel",
"parallel_chunk_indices",
]
+298 -57
View File
@@ -7,6 +7,7 @@ This module intentionally uses only PyTorch and FastVideo configuration types.
"""
import math
from collections.abc import Iterator
from dataclasses import dataclass
import torch
@@ -14,7 +15,10 @@ import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from fastvideo.attention import get_attn_backend
from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
class DiagonalGaussianDistribution:
@@ -291,6 +295,7 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
class MiniMaxH3VideoAttention(nn.Module):
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
"""Build projections and the selected dense FastVideo attention implementation."""
super().__init__()
self.heads = heads
self.dim_head = dim_head
@@ -302,12 +307,34 @@ class MiniMaxH3VideoAttention(nn.Module):
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
self.attn_impl = None
from fastvideo.platforms import current_platform
if current_platform.is_cuda_alike():
attention_backend = get_attn_backend(
dim_head,
# FlashAttention executes the FP32 VAE activations in BF16 and
# restores FP32 output, so resolve against the kernel dtype.
torch.bfloat16,
supported_attention_backends=(
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
),
)
self.attn_impl = attention_backend.get_impl_cls()(
num_heads=heads,
head_size=dim_head,
softmax_scale=dim_head**-0.5,
num_kv_heads=heads,
causal=False,
)
def forward(
self,
hidden_states: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""Apply dense self-attention to one spatial VAE token sequence."""
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
@@ -328,9 +355,17 @@ class MiniMaxH3VideoAttention(nn.Module):
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
if self.attn_impl is not None and query.device.type != "cpu":
# VAE decoding has no diffusion-step metadata, so call the selected
# backend implementation directly with dense BSHD tensors.
hidden_states = self.attn_impl.forward(query, key, value, None)
hidden_states = hidden_states.flatten(2, 3)
else:
# Keep CPU construction and execution available without requiring
# an accelerator attention backend.
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
return self.to_out[0](hidden_states)
@@ -433,6 +468,7 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Decode one latent spatial input through the H3 video transformer."""
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).reshape(
batch_size,
@@ -482,6 +518,11 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
)
def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the video decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3VideoViTDecoder3d)
class AutoencoderKLMiniMaxH3(nn.Module):
"""MiniMax-H3 causal encoder and ViT decoder with exact release geometry."""
@@ -489,6 +530,7 @@ class AutoencoderKLMiniMaxH3(nn.Module):
_no_split_modules = ["MiniMaxH3VideoResnetBlock3d", "MiniMaxH3VideoTransformerBlock"]
_repeated_blocks = ["MiniMaxH3VideoTransformerBlock"]
_keep_in_fp32_modules = ["encoder", "decoder", "quant_conv", "post_quant_conv"]
_compile_conditions = [_is_minimax_h3_video_vae_decoder]
def __init__(self, config: MiniMaxH3VideoVAEConfig) -> None:
super().__init__()
@@ -654,12 +696,15 @@ class AutoencoderKLMiniMaxH3(nn.Module):
slice_rest[dim] = slice(blend_extent, None)
return torch.cat([blended, b[tuple(slice_rest)]], dim=dim)
# The fixed spatial tile grid reuses one compiled blend-and-concatenate graph.
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
def _stitch_tiles(
self,
tiles: list[list[torch.Tensor]],
height_overlaps: list[int],
width_overlaps: list[int],
) -> torch.Tensor:
"""Blend decoded tile overlaps and concatenate the spatial canvas."""
result_rows = []
for row_index, row in enumerate(tiles):
result_row = []
@@ -676,6 +721,12 @@ class AutoencoderKLMiniMaxH3(nn.Module):
result_rows.append(torch.cat(result_row, dim=-1))
return torch.cat(result_rows, dim=-2)
# Each fixed-shape latent tile reuses one compiled decoder-input projection.
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
def _project_decoder_tile(self, tile: torch.Tensor) -> torch.Tensor:
"""Project one spatial latent tile into the decoder input channels."""
return self.post_quant_conv(tile)
def _encode_clip(self, x: torch.Tensor) -> torch.Tensor:
if not self.use_tiling:
return self.quant_conv(self.encoder(x))
@@ -699,36 +750,64 @@ class AutoencoderKLMiniMaxH3(nn.Module):
rows.append(row)
latent_y_overlaps = [overlap // self.spatial_compression_ratio for overlap in y_overlaps]
latent_x_overlaps = [overlap // self.spatial_compression_ratio for overlap in x_overlaps]
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps)
# Under mode="reduce-overhead" the stitched canvas is a CUDA-graph
# static buffer that the next _stitch_tiles replay overwrites. Callers
# (_encode/_encode_pixels/encode_keyframe) collect per-clip results
# across replays before concatenating, so hand them a caller-owned
# tensor instead of cudagraph-pooled storage.
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps).clone()
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
if not self.use_tiling:
return self.decoder(self.post_quant_conv(z))
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
for y_position, y_length in zip(y_indices, y_lengths):
row = []
for x_position, x_length in zip(x_indices, x_lengths):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
row.append(self.decoder(self.post_quant_conv(tile)))
rows.append(row)
return self._stitch_tiles(rows, y_overlaps, x_overlaps)
"""Decode one temporal clip, with optional overlapping spatial tiles."""
with nvtx_range("minimax_h3.vae.decode_clip"):
if not self.use_tiling:
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"):
projected_clip = self.post_quant_conv(z)
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"):
return self.decoder(projected_clip)
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
with nvtx_range("minimax_h3.vae.decode_clip.split_tiles"):
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
# The eager tile driver owns NVTX so each marker remains outside
# the compiled decoder graph.
with nvtx_range("minimax_h3.vae.decode_clip.decode_tiles"):
for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths)):
row = []
for column_index, (x_position, x_length) in enumerate(zip(x_indices, x_lengths)):
with nvtx_range(f"minimax_h3.vae.decode_clip.tile.{row_index}.{column_index}"):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
projected_tile = self._project_decoder_tile(tile)
with nvtx_range("minimax_h3.vae.decode_clip.tile.decoder_forward"):
decoded_tile = self.decoder(projected_tile)
row.append(decoded_tile)
rows.append(row)
with nvtx_range("minimax_h3.vae.decode_clip.stitch_tiles"):
# Same CUDA-graph output-ownership contract as _encode_clip:
# _decode collects chunks across _stitch_tiles replays before
# torch.cat, so the pooled canvas must not escape this driver.
# (The streaming _decode_to_pixels path copies each chunk out
# before the next decode and never held stale storage; the
# clone keeps that path correct too at one D2D copy per chunk.)
return self._stitch_tiles(rows, y_overlaps, x_overlaps).clone()
def _encode(self, x: torch.Tensor) -> torch.Tensor:
clip_length = self.config.clip_length
@@ -747,43 +826,157 @@ class AutoencoderKLMiniMaxH3(nn.Module):
moments = moments[:, :, :-self.config.token_drop]
return moments
def _decode(self, z: torch.Tensor) -> torch.Tensor:
tokens_chunk_size = self.tokens_chunk_size
def _encode_pixels(self, pixels: torch.Tensor) -> torch.Tensor:
"""Encode unnormalized pixels while keeping full videos off the accelerator."""
clip_length = self.config.clip_length
moments = []
for frame_start in range(0, pixels.shape[2], clip_length):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=self.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = self.normalize_pixels(clip)
moments.append(self._encode_clip(clip))
del clip
encoded = torch.cat(moments, dim=2)
if self.config.token_drop > 0:
encoded = encoded[:, :, :-self.config.token_drop]
return encoded
def _temporal_decode_plan(self, latent_num_frames: int) -> tuple[int, int, int]:
"""Return pad tokens, chunk count, and exact decoded frame count."""
if latent_num_frames <= 0:
raise ValueError(f"MiniMax-H3 latent frame count must be positive, got {latent_num_frames}.")
token_drop = self.config.token_drop
tokens_chunk_size = self.tokens_chunk_size
temporal_ratio = self.temporal_compression_ratio
chunk_num_frames = tokens_chunk_size * temporal_ratio
num_tokens = z.shape[2] + token_drop
num_tokens = latent_num_frames + token_drop
pad_tokens = (-num_tokens) % tokens_chunk_size
num_chunks = (num_tokens + pad_tokens) // tokens_chunk_size - int(token_drop > 0)
if num_chunks < 1:
pad_tokens += tokens_chunk_size
num_chunks = 1
decoded_num_frames = num_chunks * (tokens_chunk_size * temporal_ratio - self.frame_pre_padding)
if token_drop > 0:
decoded_num_frames += self.frame_overlap
if pad_tokens > 0:
intra_tail = self.config.clip_length % temporal_ratio
pad_frames = sum(intra_tail if intra_tail and (latent_num_frames + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded_num_frames -= pad_frames
if decoded_num_frames <= 0:
raise RuntimeError(
f"MiniMax-H3 decode plan produced {decoded_num_frames} frames for {latent_num_frames} latent "
"frames; the clip_length/token_drop configuration is inconsistent.")
return pad_tokens, num_chunks, decoded_num_frames
def _decode_chunks(self, z: torch.Tensor) -> Iterator[torch.Tensor]:
"""Yield finalized temporal chunks in decode order."""
tokens_chunk_size = self.tokens_chunk_size
chunk_num_frames = tokens_chunk_size * self.temporal_compression_ratio
pad_tokens, num_chunks, output_num_frames = self._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
decoded_chunks = []
output_frame_start = 0
overlap = None
for index in range(num_chunks):
start = index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
for overlap_index in range(int(token_drop > 0) + 1):
frame_start = overlap_index * chunk_num_frames
chunk = clip[:, :, frame_start:frame_start + chunk_num_frames]
chunk = chunk[:, :, self.frame_pre_padding:]
if overlap_index == 0:
for chunk_index in range(num_chunks):
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}"):
start = chunk_index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.0"):
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
if overlap is not None:
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
decoded_chunks.append(chunk)
else:
overlap = chunk
if overlap is not None:
decoded_chunks.append(overlap)
decoded = torch.cat(decoded_chunks, dim=2)
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
chunk = chunk[:, :, :num_frames]
if pad_tokens > 0:
intra_tail = self.config.clip_length % temporal_ratio
num_tokens_before_pad = z.shape[2] - pad_tokens
pad_frames = sum(intra_tail if intra_tail and (num_tokens_before_pad + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded = decoded[:, :, :-pad_frames]
return decoded
next_overlap = None
if self.config.token_drop > 0:
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.1"):
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
# Yield after the ranges close so consumer-side CPU copies do not inflate decoder timing.
overlap = next_overlap
if num_frames > 0:
output_frame_start += num_frames
yield chunk
if overlap is not None and output_frame_start < output_num_frames:
yield overlap[:, :, :output_num_frames - output_frame_start]
def decoded_pixel_shape(self, latent_shape: torch.Size | tuple[int, ...]) -> tuple[int, int, int, int, int]:
"""Return the exact CPU pixel-buffer shape for a latent tensor shape."""
if len(latent_shape) != 5:
raise ValueError(f"MiniMax-H3 latents must be five-dimensional, got shape {tuple(latent_shape)}.")
batch_size, channels, latent_num_frames, latent_height, latent_width = map(int, latent_shape)
if channels != self.latent_channels:
raise ValueError(f"MiniMax-H3 latents must have {self.latent_channels} channels, got {channels}.")
_, _, decoded_num_frames = self._temporal_decode_plan(latent_num_frames)
return (
batch_size,
int(self.config.out_channels),
decoded_num_frames,
latent_height * self.spatial_compression_ratio,
latent_width * self.spatial_compression_ratio,
)
@staticmethod
def _streams_chunk_copies(z: torch.Tensor, output: torch.Tensor) -> bool:
"""Whether finalized chunks copy to ``output`` asynchronously on the current CUDA stream."""
return z.device.type == "cuda" and output.is_pinned()
@staticmethod
def _copy_chunk_pixels(pixels: torch.Tensor, output: torch.Tensor, frame_start: int, non_blocking: bool) -> None:
"""Copy one finalized fp32 pixel chunk into the CPU ``output`` buffer.
Device-to-host copies run per (batch, channel) plane: the temporal
slice of ``output`` is strided across channels, but each plane is
contiguous on both sides, so every transfer stays a direct memcpy
instead of staging through a pageable CPU temporary. With a pinned
``output`` and ``non_blocking=True`` the copies are additionally
asynchronous on the current CUDA stream; callers synchronize once
before releasing the buffer.
"""
target = output[:, :, frame_start:frame_start + pixels.shape[2]]
if pixels.device.type == "cuda":
pixels = pixels.contiguous()
for batch_index in range(pixels.shape[0]):
for channel_index in range(pixels.shape[1]):
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
non_blocking=non_blocking)
else:
target.copy_(pixels)
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
Each finalized chunk streams through ``_copy_chunk_pixels`` (direct
per-plane memcpys; asynchronous with a pinned ``output``) so the
copies overlap the next chunk's decode; ``decode_to_pixels``
synchronizes once before returning.
"""
non_blocking = self._streams_chunk_copies(z, output)
output_frame_start = 0
for chunk in self._decode_chunks(z):
num_frames = chunk.shape[2]
pixels = self.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._copy_chunk_pixels(pixels, output, output_frame_start, non_blocking)
output_frame_start += num_frames
if output_frame_start != output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {output_frame_start} frames into an output buffer expecting "
f"{output.shape[2]}.")
def _decode(self, z: torch.Tensor) -> torch.Tensor:
return torch.cat(list(self._decode_chunks(z)), dim=2)
def encode(
self,
@@ -799,6 +992,34 @@ class AutoencoderKLMiniMaxH3(nn.Module):
return (posterior, )
return AutoencoderKLOutput(latent_dist=posterior)
def encode_pixels(
self,
pixels: torch.Tensor,
return_dict: bool = True,
) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]:
"""Encode CPU-resident pixels one VAE clip at a time.
``pixels`` stays on CPU as ``uint8`` in ``[0, 255]`` or floating point
in ``[0, 1]``; each clip is moved to the VAE device, normalized, and
encoded so only one clip of pixels is resident on the accelerator.
"""
if pixels.ndim != 5 or pixels.shape[1] != self.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {self.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if self.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([self._encode_pixels(pixel_slice) for pixel_slice in pixels.split(1)])
else:
moments = self._encode_pixels(pixels)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior, )
return AutoencoderKLOutput(latent_dist=posterior)
def encode_keyframe(
self,
x: torch.Tensor,
@@ -825,6 +1046,26 @@ class AutoencoderKLMiniMaxH3(nn.Module):
return (decoded, )
return DecoderOutput(sample=decoded)
def decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
"""Stream decoded ``[0, 1]`` FP32 pixels into a caller-owned CPU buffer."""
expected_shape = self.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
try:
if self.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
self._decode_to_pixels(z_slice, output[batch_index:batch_index + 1])
else:
self._decode_to_pixels(z, output)
finally:
# Drain async chunk copies before the caller (or an exception
# handler) can read or release the pinned buffer.
if self._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def forward(
self,
sample: torch.Tensor,
@@ -12,9 +12,9 @@ from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_IMAGE_PAD_TOKEN,
MINIMAX_H3_TEXT_ENCODER_LAYER,
MINIMAX_H3_TEXT_TAG,
MINIMAX_H3_VIDEO_PAD_TOKEN,
MINIMAX_H3_VIDEO_TAG,
@@ -42,25 +42,6 @@ def _token_ids(tokenized: Any) -> list[int]:
return [int(token_id) for token_id in input_ids]
def _create_mm_token_type_ids(processor: Any, token_ids: list[int]) -> list[list[int]]:
"""Build Qwen3-VL modality IDs across old and new Transformers releases."""
create_ids = getattr(processor, "create_mm_token_type_ids", None)
if callable(create_ids):
return create_ids([token_ids])
modality_ids = [0] * len(token_ids)
for modality, modality_type in (("image", 1), ("video", 2), ("audio", 3)):
special_ids = getattr(processor, f"{modality}_token_ids", None)
if special_ids is None:
special_id = getattr(processor, f"{modality}_token_id", None)
special_ids = [] if special_id is None else [special_id]
resolved_ids = {int(special_id) for special_id in special_ids if special_id is not None}
for index, token_id in enumerate(token_ids):
if token_id in resolved_ids:
modality_ids[index] = modality_type
return [modality_ids]
def build_ref2va_presentation(
tokenizer: Any,
prompt: str,
@@ -155,20 +136,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
device: torch.device,
**vision_inputs: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
hidden_state_index = MINIMAX_H3_TEXT_ENCODER_LAYER
input_ids = torch.tensor([token_ids], dtype=torch.long, device=device)
mm_token_type_ids = torch.as_tensor(
_create_mm_token_type_ids(self.processor, token_ids),
dtype=torch.long,
device=device,
)
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
dtype = self.conditioner.dtype
outputs = self.conditioner(
input_ids=input_ids,
attention_mask=torch.ones_like(input_ids),
mm_token_type_ids=mm_token_type_ids,
use_cache=False,
output_hidden_states=True,
prompt_embeds = self.conditioner(
input_ids,
**{
name:
None if value is None else value.to(
@@ -178,10 +149,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
for name, value in vision_inputs.items()
},
)
if outputs.hidden_states is None or len(outputs.hidden_states) <= hidden_state_index:
raise ValueError(f"Qwen3-VL did not return `hidden_states[{hidden_state_index}]`.")
if prompt_embeds.ndim != 2 or prompt_embeds.shape[0] != len(token_ids):
raise ValueError(f"MiniMax-H3 slim text encoder returned unexpected shape={tuple(prompt_embeds.shape)}")
return (
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
torch.tensor(token_tags, dtype=torch.long),
)
@@ -286,6 +257,7 @@ class MiniMaxH3ConditioningStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Encode one H3 prompt presentation and attach its packed text features."""
device = get_local_torch_device()
first_param = next(self.conditioner.parameters(), None)
moved_for_forward = (fastvideo_args.text_encoder_cpu_offload and first_param is not None
@@ -293,10 +265,13 @@ class MiniMaxH3ConditioningStage(PipelineStage):
if moved_for_forward:
self.conditioner.to(device)
try:
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
# Keep both H3 prompt-presentation modes under one text-encoding
# range so Nsight Systems exposes their complete conditioning cost.
with nvtx_range("minimax_h3.text_encoding"):
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
finally:
if moved_for_forward:
self.conditioner.to("cpu")
@@ -7,10 +7,13 @@ from typing import Any
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE
from fastvideo.models.vaes.minimax_h3_parallel import DEFAULT_DECODE_GATHER_STRATEGY, decode_to_pixels_parallel
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MiniMaxH3PackedLayout,
unpack_audio_tokens,
@@ -21,6 +24,9 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.utils import is_pin_memory_available
logger = init_logger(__name__)
def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
@@ -30,6 +36,23 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
return layout
def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]:
"""Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages.
The executors consume rank 0's ForwardBatch and the training validation
callback consumes each sequence-parallel group leader's, so the output
rank is the SP group's first rank (identical to world rank 0 in the
single-group e2e case). ``parallel`` is only true when every group rank
will run the decode body — the collectives inside require uniform
participation, so no rank-dependent branch may guard them.
"""
if not model_parallel_is_initialized():
return None, True, False
sp_group = get_sp_group()
parallel = bool(want_parallel) and sp_group.world_size > 1
return sp_group, sp_group.is_first_rank, parallel
class MiniMaxH3VideoDecodingStage(PipelineStage):
"""Drop visual condition rows, unpatchify, and decode the target video."""
@@ -54,6 +77,16 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 video latents into normalized CPU pixels."""
placeholder = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
sp_group, is_output_rank, parallel = _decode_participation(fastvideo_args, fastvideo_args.vae_parallel_decode)
if not is_output_rank and not parallel:
# Consumers read the output rank's ForwardBatch. Keep a
# verifier-compatible placeholder on other ranks and avoid
# duplicating the full VAE decode and CPU output buffer.
batch.output = placeholder
return batch
layout = _layout(batch)
if batch.latents is None or batch.raw_latent_shape is None or len(batch.raw_latent_shape) != 5:
raise ValueError("MiniMax-H3 video latents or raw geometry are missing at decode.")
@@ -71,13 +104,33 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
try:
latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32))
if fastvideo_args.output_type == "latent":
batch.output = latents.detach().float().cpu()
# No collectives on this path, so uniform participation is
# trivial: every rank returns here.
batch.output = latents.detach().float().cpu() if is_output_rank else placeholder
return batch
# The published decode recipe uses FP16 autocast over FP32 weights.
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
video = self.vae.decode(latents).sample
batch.output = self.vae.denormalize_pixels(video.float()).clamp_(0, 1).cpu()
output = None
if is_output_rank:
output = torch.empty(
self.vae.decoded_pixel_shape(latents.shape),
device="cpu",
dtype=torch.float32,
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
)
# Attribute the streamed decoder computation while retaining
# per-chunk device-to-host transfer and pinned-buffer reuse.
with (
nvtx_range("minimax_h3.vae"),
torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"),
):
if parallel:
strategy = fastvideo_args.vae_parallel_decode_strategy or DEFAULT_DECODE_GATHER_STRATEGY
logger.info_once(f"MiniMax-H3 VAE decode: sequence-parallel chunks across "
f"{sp_group.world_size} ranks ({strategy})")
decode_to_pixels_parallel(self.vae, latents, output, sp_group, strategy=strategy)
else:
self.vae.decode_to_pixels(latents, output)
batch.output = output if is_output_rank else placeholder
return batch
finally:
if fastvideo_args.vae_cpu_offload:
@@ -107,6 +160,15 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 audio latents into a stereo CPU waveform."""
# Audio decode is sub-second, so it always runs serially on the SP
# group's first rank (the rank whose ForwardBatch consumers read).
if model_parallel_is_initialized() and not get_sp_group().is_first_rank:
batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32)
batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate
self._clear_runtime(batch)
return batch
layout = _layout(batch)
if batch.audio_latents is None:
raise ValueError("MiniMax-H3 audio latents are missing at decode.")
@@ -124,7 +186,10 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
self._clear_runtime(batch)
return batch
decoded = self.audio_vae.decode(latents).sample.float()
# The range isolates waveform synthesis from packing and runtime
# cleanup so the audio decoder has one stable timeline boundary.
with nvtx_range("minimax_h3.audio_vae"):
decoded = self.audio_vae.decode(latents).sample.float()
if decoded.ndim != 3 or decoded.shape[0] != 2 or decoded.shape[1] != 1:
raise ValueError("MiniMax-H3 audio VAE must decode stereo channels as two mono batch items; "
f"got {tuple(decoded.shape)}.")
@@ -11,8 +11,8 @@ from fastvideo.attention.selector import component_attention_backend, get_attn_b
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.profiler import profiler_region
from fastvideo.hooks.activation_trace import trace_step
from fastvideo.profiler import nvtx_range, profiler_region
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_KEYFRAME_NOISE_AUG,
MiniMaxH3PackedLayout,
@@ -89,6 +89,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Denoise the packed H3 video and audio streams over one shared schedule."""
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
if not isinstance(layout, MiniMaxH3PackedLayout):
raise ValueError("MiniMax-H3 packed layout is missing before denoising.")
@@ -145,9 +146,15 @@ class MiniMaxH3DenoisingStage(PipelineStage):
vsa_exempt = vsa_mode == "exempt"
vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ()))
vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0))
# Run-level tile geometry (256 default, 64 = native Triton path),
# plumbed like the run-level sparsity; the builder validates the
# value against VSA_H3_TILE_SHAPES.
vsa_tile_size = int(fastvideo_args.VSA_tile_size)
try:
with profiler_region("inference_denoising"):
# The stage range groups the complete denoising loop while the
# indexed model ranges retain timing detail for every H3 block.
with profiler_region("inference_denoising"), nvtx_range("minimax_h3.dit"):
for index, (video_timestep,
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
unique_timesteps, timestep_indices = row_timestep_plan[index]
@@ -167,6 +174,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
device=device,
exempt=vsa_exempt,
dense_layers=vsa_dense_layers,
tile_size=vsa_tile_size,
)
# Under torch.compile(mode="reduce-overhead") each denoising
# step must be marked, or cudagraph trees flag cross-step
@@ -9,8 +9,10 @@ import numpy as np
import torch
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vaes.minimax_h3_parallel import encode_pixels_parallel
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_AUDIO_CHANNELS,
MINIMAX_H3_KEYFRAME_ENCODE_SEED,
@@ -36,6 +38,8 @@ from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
MINIMAX_H3_LAYOUT_KEY = "minimax_h3_layout"
@@ -105,8 +109,20 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
self,
references: list[MiniMaxH3PreparedReference],
device: torch.device,
fastvideo_args: FastVideoArgs,
) -> list[torch.Tensor]:
patch_size = self.transformer.patch_size
# Reference encode runs on every rank (all ranks hold identical
# prepared references), so clip-parallel encode keeps participation
# uniform by construction: each rank encodes a clip subset and the
# all-gather leaves the identical full posterior everywhere.
parallel_group = None
if fastvideo_args.vae_parallel_encode and model_parallel_is_initialized():
sp_group = get_sp_group()
if sp_group.world_size > 1:
parallel_group = sp_group
logger.info_once(f"MiniMax-H3 reference VAE encode: sequence-parallel clips across "
f"{sp_group.world_size} ranks")
rows: list[torch.Tensor] = []
for reference in references:
if reference.media_type == "audio":
@@ -119,9 +135,11 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
if reference.frames is None:
raise ValueError("MiniMax-H3 reference video frames are missing.")
frames = reference.frames[:trim_reference_num_frames(reference.frames.shape[0])]
pixels = torch.from_numpy(frames.copy()).permute(3, 0, 1, 2)[None]
pixels = pixels.to(device=device, dtype=torch.float32).div_(255.0)
posterior = self.vae.encode(self.vae.normalize_pixels(pixels)).latent_dist
pixels = torch.from_numpy(np.ascontiguousarray(frames)).permute(3, 0, 1, 2)[None]
if parallel_group is not None:
posterior = encode_pixels_parallel(self.vae, pixels, parallel_group).latent_dist
else:
posterior = self.vae.encode_pixels(pixels).latent_dist
latents = self.vae.normalize_latents(_sample_visual_posterior(posterior).to(
torch.float16).float()).cpu()
reference.num_latent_frames = int(latents.shape[2])
@@ -202,7 +220,7 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
vae_device = get_local_torch_device()
self.vae.to(vae_device)
try:
video_rows = self._encode_visual_rows(references, vae_device)
video_rows = self._encode_visual_rows(references, vae_device, fastvideo_args)
finally:
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
+4 -1
View File
@@ -48,7 +48,10 @@ class MpsPlatform(Platform):
@classmethod
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None, head_size: int,
dtype: torch.dtype) -> str:
# MPS supports SDPA (Scaled Dot-Product Attention) which is the most compatible
if selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
raise NotImplementedError("VIDEO_SPARSE_ATTN is not supported on MPS. Unset "
"FASTVIDEO_ATTENTION_BACKEND or set it to TORCH_SDPA.")
# MPS supports SDPA (Scaled Dot-Product Attention) which is the most compatible.
logger.info("Using Torch SDPA backend for MPS.")
return "fastvideo.attention.backends.sdpa.SDPABackend"
+19
View File
@@ -38,6 +38,25 @@ logger = init_logger(__name__)
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
@contextlib.contextmanager
def nvtx_range(name: str):
"""Emit one optional NVTX range for an external CUDA profiler.
``FASTVIDEO_NVTX_PROFILE=1`` enables the marker. The context manager stays
a no-op without CUDA so call sites can remain shared with CPU tests.
"""
enabled = envs.FASTVIDEO_NVTX_PROFILE and torch.cuda.is_available()
if not enabled:
yield
return
torch.cuda.nvtx.range_push(name)
try:
yield
finally:
torch.cuda.nvtx.range_pop()
@dataclass(frozen=True)
class ProfilerRegion:
"""Metadata describing a profiler region."""
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU backward checks for the VSA-H3 backend.
The CuTe backend returns FA4's own output tensor, which FA4's autograd node
saved for its backward. Composing the compression branch onto it in place
therefore poisons the graph, and the failure only appears once the VSA-256
CuTe path has a backward at all. These tests pin the composition.
"""
import pytest
import torch
from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAImpl, MiniMaxH3VSAMetadataBuilder)
_SPEC = dict(raw_latent_shape=(16, 16, 24), patch_size=(1, 2, 2), prefix_segments=(64, 32, 16))
_HEADS = 2
_DIM = 128
def _build_meta(device, sparsity=0.5):
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=_SPEC["raw_latent_shape"],
patch_size=_SPEC["patch_size"],
VSA_sparsity=sparsity,
prefix_segments=_SPEC["prefix_segments"],
device=device,
)
def _select_backend(monkeypatch, backend):
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _forward_backward(impl, meta, gate_compress, device):
seq = meta.total_seq_length
torch.manual_seed(0)
q, k, v = (torch.randn(1, seq, _HEADS, _DIM, device=device, dtype=torch.bfloat16, requires_grad=True)
for _ in range(3))
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
gate = None
if gate_compress:
gate = torch.randn(1, tq.shape[1], _HEADS, _DIM, device=device, dtype=torch.bfloat16) * 0.1
out = impl.forward(tq, tk, tv, gate, meta)
out = impl.postprocess_output(out, meta)
out.float().pow(2).sum().backward()
return out, (q, k, v)
@pytest.mark.parametrize("backend", ["triton", "cute"])
@pytest.mark.parametrize("gate_compress", [False, True])
def test_h3_vsa_backward_runs(monkeypatch, backend: str, gate_compress: bool) -> None:
"""Regression: with the CuTe backend and a non-zero gate this used to die
with "one of the variables needed for gradient computation has been
modified by an inplace operation ... output 0 of FlashAttnFuncBackward".
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
_select_backend(monkeypatch, backend)
device = torch.device("cuda")
meta = _build_meta(device)
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
out, leaves = _forward_backward(impl, meta, gate_compress, device)
assert torch.isfinite(out).all().item()
for name, leaf in zip(("q", "k", "v"), leaves):
assert leaf.grad is not None, f"{name} received no gradient"
assert torch.isfinite(leaf.grad).all().item(), f"{name}.grad has non-finite values"
assert leaf.grad.abs().sum().item() > 0, f"{name}.grad is all zero"
@pytest.mark.parametrize("gate_compress", [False, True])
def test_h3_vsa_backward_cute_matches_triton(monkeypatch, gate_compress: bool) -> None:
"""CuTe and Triton take different routes to the same math; their gradients
should agree to bf16 tolerance."""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
device = torch.device("cuda")
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
grads = {}
for backend in ("triton", "cute"):
with monkeypatch.context() as m:
_select_backend(m, backend)
meta = _build_meta(device)
_, leaves = _forward_backward(impl, meta, gate_compress, device)
grads[backend] = [leaf.grad.detach().float() for leaf in leaves]
for name, ref, got in zip(("dq", "dk", "dv"), grads["triton"], grads["cute"]):
diff = (ref - got).abs()
avg_abs = diff.mean().item()
max_rel = (diff.max() / (ref.abs().mean() + 1e-6)).item()
print(f"[h3-vsa gate={gate_compress}] {name}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < 1e-2, f"{name}: avg_abs {avg_abs:.3e}"
assert max_rel < 0.5, f"{name}: max_rel {max_rel:.3e}"
@@ -5,20 +5,29 @@ reference. The same reference doubles as the GPU kernel parity oracle."""
import math
import pytest
import torch
import torch.nn.functional as F
from fastvideo.attention.backends.video_sparse_attn_h3 import (_TILE_ELEMS, MiniMaxH3VSAImpl,
MiniMaxH3VSAMetadataBuilder, _build_block_mask,
_pool_tiles, token_tile_and_valid)
_pool_tiles, _validate_h3_tile_geometry,
token_tile_and_valid)
_720P = dict(raw_latent_shape=(30, 44, 80), patch_size=(1, 2, 2), prefix_segments=(512, 1760, 400))
_TINY = dict(raw_latent_shape=(8, 8, 12), patch_size=(1, 2, 2), prefix_segments=(7, 5, 3))
# (4,4,4) coverage: dit grid (9, 10, 13) is ragged in all three dims
# (t: 4+4+1, h: 4+4+2, w: 4+4+4+1) and every prefix segment leaves a
# partial tail tile at 64 (70 -> 64+6, 5 -> 5, 130 -> 64+64+2).
_TINY64 = dict(raw_latent_shape=(9, 20, 26), patch_size=(1, 2, 2), prefix_segments=(70, 5, 130))
# production-shape request: 768x1344, 124 frames -> latents (37, 48, 84),
# patch (1,2,2) -> token grid (37, 24, 42); text 300 + audio 414 rows.
_PROD = dict(raw_latent_shape=(37, 48, 84), patch_size=(1, 2, 2), prefix_segments=(300, 0, 414))
_CPU = torch.device("cpu")
def _build(spec, sparsity=0.0, device=_CPU):
def _build(spec, sparsity=0.0, device=_CPU, tile_size=_TILE_ELEMS):
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=spec["raw_latent_shape"],
@@ -26,6 +35,7 @@ def _build(spec, sparsity=0.0, device=_CPU):
VSA_sparsity=sparsity,
prefix_segments=spec["prefix_segments"],
device=device,
tile_size=tile_size,
)
@@ -36,7 +46,7 @@ def _impl():
def reference_sparse_attention(query, key, value, mask, meta):
"""Token-level oracle: SDPA over the padded tile buffer with the block
mask expanded to tokens. query/key/value: tiled [B, S_pad, H, D]."""
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes)
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes, meta.tile_elems)
out = torch.empty_like(query)
for b in range(query.shape[0]):
for h in range(query.shape[2]):
@@ -133,9 +143,117 @@ def test_prefix_queries_stay_dense_at_high_sparsity():
"video rows should actually be sparse at 75%"
# ---------------------------------------------------------------------------
# 64-token (4,4,4) tile geometry
# ---------------------------------------------------------------------------
def test_geometry_tile64_ragged_tails():
"""Hand-computed (4,4,4) oracle on a grid ragged in all three dims."""
meta = _build(_TINY64, tile_size=64)
assert meta.tile_elems == 64
t, h, w = 9, 10, 13 # raw latents (9, 20, 26) under patch (1, 2, 2)
n_t, n_h, n_w = 3, 3, 4
prefix_len = sum(_TINY64["prefix_segments"])
seq = prefix_len + t * h * w
assert meta.total_seq_length == seq
assert meta.num_prefix_tiles == 2 + 1 + 3
assert meta.num_video_tiles == n_t * n_h * n_w
assert int(meta.variable_block_sizes.sum()) == seq
assert int(meta.variable_block_sizes.max()) <= 64
assert meta.variable_block_sizes[:meta.num_prefix_tiles].tolist() == [64, 6, 5, 64, 64, 2]
# per-tile valid sizes: product of the per-dim clamped tails
expected = torch.tensor([
min(4, t - 4 * tt) * min(4, h - 4 * hh) * min(4, w - 4 * ww) for tt in range(n_t) for hh in range(n_h)
for ww in range(n_w)
],
dtype=torch.long)
assert torch.equal(meta.variable_block_sizes[meta.num_prefix_tiles:], expected)
assert int(expected.min()) == 1 * 2 * 1 # the (t,h,w) ragged corner
# every packed video row lands in the 3D tile its (t,h,w) coordinate says
idx = meta.untile_combined_index
row = torch.arange(t * h * w)
row_t, row_h, row_w = row // (h * w), (row // w) % h, row % w
expected_tile = meta.num_prefix_tiles + ((row_t // 4) * n_h + row_h // 4) * n_w + row_w // 4
assert torch.equal(idx[prefix_len:] // 64, expected_tile)
# and in a non-pad slot of that tile
assert bool((idx % 64 < meta.variable_block_sizes[idx // 64]).all())
# untile(tile(x)) == x on the 64-wide padded buffer
x = torch.randn(1, seq, 2, 4)
buf = _impl().tile(x, meta)
assert buf.shape[1] == meta.variable_block_sizes.numel() * 64
assert torch.equal(buf[:, idx], x)
def test_geometry_tile64_production_shape():
"""Production latents (37, 48, 84): ragged t and w tails at (4,4,4)."""
meta64 = _build(_PROD, tile_size=64)
assert meta64.num_prefix_tiles == 5 + 7 # 300 -> 4x64+44, 414 -> 6x64+30
assert meta64.num_video_tiles == 10 * 6 * 11 # (37, 24, 42) / (4, 4, 4)
assert meta64.total_seq_length == 300 + 414 + 37 * 24 * 42
assert int(meta64.variable_block_sizes.sum()) == meta64.total_seq_length
sizes_vid = meta64.variable_block_sizes[meta64.num_prefix_tiles:]
assert int(sizes_vid.max()) == 64 and int(sizes_vid.min()) == 1 * 4 * 2 # (t, w) ragged corner
# same packed sequence under the default 256 geometry, fewer tiles
meta256 = _build(_PROD)
assert meta256.tile_elems == _TILE_ELEMS
assert meta256.num_prefix_tiles == 2 + 2
assert meta256.num_video_tiles == 10 * 3 * 6
assert meta256.total_seq_length == meta64.total_seq_length
x = torch.randn(1, meta64.total_seq_length, 2, 4)
buf = _impl().tile(x, meta64)
assert torch.equal(buf[:, meta64.untile_combined_index], x)
def test_sparsity_zero_matches_dense_sdpa_tile64():
torch.manual_seed(2)
meta = _build(_TINY64, tile_size=64)
seq = meta.total_seq_length
q, k, v = (torch.randn(1, seq, 2, 8) for _ in range(3))
impl = _impl()
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
scores = torch.matmul(_pool_tiles(tq, meta.variable_block_sizes, meta.tile_elems),
_pool_tiles(tk, meta.variable_block_sizes, meta.tile_elems).transpose(-2, -1))
mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, 0.0, exempt=True)
sparse_out = impl.postprocess_output(reference_sparse_attention(tq, tk, tv, mask, meta), meta)
dense_out = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2)
assert torch.allclose(sparse_out, dense_out, atol=1e-5), (sparse_out - dense_out).abs().max()
def test_geometry_guard_enforces_tile64_bound():
"""A 65-token tile passes the 256 bound but must fail the 64 one."""
meta = _build(_TINY64, tile_size=64)
prefix = tuple(s for s in _TINY64["prefix_segments"] if s > 0)
dit_shape = (9, 10, 13)
sizes = meta.variable_block_sizes.clone()
sizes[0] = 65
with pytest.raises(ValueError, match="tile sizes out of bounds"):
_validate_h3_tile_geometry(prefix, dit_shape, sizes, meta.untile_combined_index, 64)
# the untampered tile-64 geometry passes its own bound
_validate_h3_tile_geometry(prefix, dit_shape, meta.variable_block_sizes, meta.untile_combined_index, 64)
def test_builder_rejects_unknown_tile_size():
for bad in (0, 128, 512):
with pytest.raises(ValueError, match="tile_size"):
_build(_TINY, tile_size=bad)
if __name__ == "__main__":
test_geometry_720p()
test_mask_policy()
test_sparsity_zero_matches_dense_sdpa()
test_prefix_queries_stay_dense_at_high_sparsity()
test_geometry_tile64_ragged_tails()
test_geometry_tile64_production_shape()
test_sparsity_zero_matches_dense_sdpa_tile64()
test_geometry_guard_enforces_tile64_bound()
test_builder_rejects_unknown_tile_size()
print("all VSA-H3 CPU checks passed")
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU checks for the VSA-H3 tile-64 sm_100a route selection.
The opt-in third kernel route (``FASTVIDEO_VSA_SM100A=1``) must (a) stay off by
default, (b) engage only when the extension is present, the device qualifies,
and the forward carries no grad, and (c) fall back to the Triton-64 entry with
one warning when the env is set but a precondition fails. All device/extension
probes are monkeypatched; no GPU or kernel install needed.
"""
import pytest
import torch
import fastvideo.attention.backends.video_sparse_attn_h3 as vsa_h3
from fastvideo.attention.backends.video_sparse_attn_h3 import (VSA_SM100A_ENV, MiniMaxH3VSAImpl,
MiniMaxH3VSAMetadataBuilder, _sm100a_unavailable_reason)
# Small tile-64 geometry: 2 prefix segments + a (4,4,8)-token video grid.
_SPEC = dict(raw_latent_shape=(4, 8, 16), patch_size=(1, 2, 2), prefix_segments=(70, 30))
_HEADS, _DIM = 2, 128
def _build_meta():
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=_SPEC["raw_latent_shape"],
patch_size=_SPEC["patch_size"],
VSA_sparsity=0.0,
prefix_segments=_SPEC["prefix_segments"],
device=torch.device("cpu"),
tile_size=64,
)
def _tiled_qkv(meta, requires_grad=False):
# bf16 like the real tiled buffers, so forward()'s dtype-cast warning
# stays out of the warning assertions below.
s_pad = meta.variable_block_sizes.numel() * 64
return tuple(
torch.randn(1, s_pad, _HEADS, _DIM, dtype=torch.bfloat16, requires_grad=requires_grad) for _ in range(3))
class _FakeSm100a:
"""Stands in for fastvideo_kernel.block_sparse_attn_sm100a."""
def __init__(self, supported=True):
self.supported = supported
self.calls = []
def is_supported(self, q, variable_block_sizes):
return self.supported
def block_sparse_attn_sm100a(self, q, k, v, q2k_idx, q2k_num, variable_block_sizes, need_lse=True):
self.calls.append(dict(q=q, q2k_idx=q2k_idx, q2k_num=q2k_num, vbs=variable_block_sizes,
need_lse=need_lse))
return q.clone(), None
def _fake_map_to_index(block_map):
"""Pure-torch stand-in for the Triton map_to_index (same contract)."""
b, h, t, n = block_map.shape
idx = torch.full((b, h, t, n), -1, dtype=torch.int32)
num = block_map.sum(dim=-1, dtype=torch.int32)
for bi in range(b):
for hi in range(h):
for ti in range(t):
cols = torch.nonzero(block_map[bi, hi, ti], as_tuple=False).flatten()
idx[bi, hi, ti, :cols.numel()] = cols.to(torch.int32)
return idx, num
class _FakeTriton:
def __init__(self):
self.calls = 0
def __call__(self, q, k, v, mask, variable_block_sizes):
self.calls += 1
return q.clone(), None
@pytest.fixture()
def routed(monkeypatch):
"""Backend with both kernel entries faked; returns (fakes, run)."""
fake_sm = _FakeSm100a()
fake_triton = _FakeTriton()
monkeypatch.setattr(vsa_h3, "_sm100a", fake_sm)
monkeypatch.setattr(vsa_h3, "block_sparse_attn_64_bhsd", fake_triton)
monkeypatch.setattr(vsa_h3, "map_to_index", _fake_map_to_index)
meta = _build_meta()
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
def run(requires_grad=False):
q, k, v = _tiled_qkv(meta, requires_grad=requires_grad)
return impl.forward(q, k, v, None, meta)
return fake_sm, fake_triton, run, meta
def test_reason_covers_every_precondition():
q = torch.randn(1, _HEADS, 128, _DIM)
vbs = torch.full((2, ), 64, dtype=torch.long)
assert "not installed" in _sm100a_unavailable_reason(None, q, vbs, grad_mode=False)
ok = _FakeSm100a(supported=True)
assert "forward-only" in _sm100a_unavailable_reason(ok, q, vbs, grad_mode=True)
bad = _FakeSm100a(supported=False)
assert "is_supported" in _sm100a_unavailable_reason(bad, q, vbs, grad_mode=False)
assert _sm100a_unavailable_reason(ok, q, vbs, grad_mode=False) is None
def test_default_off_routes_triton(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
monkeypatch.delenv(VSA_SM100A_ENV, raising=False)
run()
assert fake_triton.calls == 1
assert fake_sm.calls == []
def test_env_on_routes_sm100a_with_index_metadata(routed, monkeypatch):
fake_sm, fake_triton, run, meta = routed
monkeypatch.setenv(VSA_SM100A_ENV, "1")
out = run()
assert fake_triton.calls == 0
assert len(fake_sm.calls) == 1
call = fake_sm.calls[0]
n_tiles = meta.variable_block_sizes.numel()
# sparsity 0 -> all-True mask -> every row's count is n_tiles
assert call["q2k_num"].dtype == torch.int32 and (call["q2k_num"] == n_tiles).all()
assert call["q2k_idx"].shape[-1] == n_tiles and call["q2k_idx"].dtype == torch.int32
assert call["vbs"].dtype == torch.int32
assert call["need_lse"] is False
# BHSD kernel result comes back in the backend's BSHD layout
assert out.shape == (1, n_tiles * 64, _HEADS, _DIM)
def test_env_on_grad_inputs_fall_back_to_triton(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
monkeypatch.setenv(VSA_SM100A_ENV, "1")
run(requires_grad=True)
assert fake_triton.calls == 1
assert fake_sm.calls == []
# ...but the same process still routes no-grad forwards to sm_100a
run(requires_grad=False)
assert len(fake_sm.calls) == 1
def test_env_on_unsupported_warns_once_and_falls_back(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
fake_sm.supported = False
monkeypatch.setenv(VSA_SM100A_ENV, "1")
warnings = []
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
run()
run()
assert fake_triton.calls == 2
assert fake_sm.calls == []
assert len(warnings) == 2 # warning_once dedups by message; both carry the same one line
assert warnings[0] == warnings[1]
assert VSA_SM100A_ENV in warnings[0] and "is_supported" in warnings[0]
def test_env_on_missing_module_warns_and_falls_back(routed, monkeypatch):
fake_sm, fake_triton, run, _ = routed
monkeypatch.setattr(vsa_h3, "_sm100a", None)
monkeypatch.setenv(VSA_SM100A_ENV, "1")
warnings = []
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
run()
assert fake_triton.calls == 1
assert warnings and "not installed" in warnings[0]
def test_env_on_no_grad_context_detaches_route_from_leaf_flags(routed, monkeypatch):
"""A requires_grad leaf under torch.no_grad() is still a no-grad forward."""
fake_sm, fake_triton, run, meta = routed
monkeypatch.setenv(VSA_SM100A_ENV, "1")
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
q, k, v = _tiled_qkv(meta, requires_grad=True)
with torch.no_grad():
impl.forward(q, k, v, None, meta)
assert len(fake_sm.calls) == 1
assert fake_triton.calls == 0
@@ -15,6 +15,12 @@ import json
import os
import subprocess
import sys
from unittest.mock import Mock
import pytest
import torch
from fastvideo.profiler import nvtx_range
# Five-window child: ops before any region, inside a region, between regions,
# inside a second (short-named) region, after the last region. Exits without
@@ -105,3 +111,73 @@ def test_noop_without_profiler_dir(tmp_path):
proc = subprocess.run([sys.executable, "-c", child], env=env,
capture_output=True, text=True, timeout=300)
assert proc.returncode == 0, proc.stderr
def test_nvtx_range_disabled_is_noop(monkeypatch):
"""Keep CUDA NVTX untouched when external profiling is disabled."""
range_push = Mock()
range_pop = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
with nvtx_range("disabled"):
body_executed = True
assert body_executed is True
range_push.assert_not_called()
range_pop.assert_not_called()
def test_nvtx_range_without_cuda_is_noop(monkeypatch):
"""Keep NVTX untouched when profiling is enabled on a CPU-only process."""
range_push = Mock()
range_pop = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
with nvtx_range("cpu-only"):
body_executed = True
assert body_executed is True
range_push.assert_not_called()
range_pop.assert_not_called()
def test_nvtx_range_enabled_orders_push_body_pop(monkeypatch):
"""Place the profiled body between one matching NVTX push and pop."""
events = []
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
with nvtx_range("minimax_h3.test"):
events.append(("body", None))
assert events == [
("push", "minimax_h3.test"),
("body", None),
("pop", None),
]
def test_nvtx_range_body_exception_pops_and_propagates(monkeypatch):
"""Balance the NVTX stack while preserving a body exception."""
events = []
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
with pytest.raises(RuntimeError, match="profile body failed"):
with nvtx_range("minimax_h3.failure"):
raise RuntimeError("profile body failed")
assert events == [
("push", "minimax_h3.failure"),
("pop", None),
]
@@ -0,0 +1,281 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
import os
import pytest
import torch
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29514")
import fastvideo.models.encoders.minimax_h3_checkpoint_fp8 as h3_fp8
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.layers.linear import ColumnParallelLinear, UnquantizedLinearMethod
from fastvideo.layers.vocab_parallel_embedding import UnquantizedEmbeddingMethod, VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import (
MiniMaxH3SerializedFP8Config,
MiniMaxH3SerializedFP8LinearMethod,
)
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_read_text_encoder_checkpoint_quantization_config,
)
def _checkpoint_quantization_config(**overrides) -> dict:
config = {
"quant_method": "fp8",
"activation_scheme": "dynamic",
"fmt": "e4m3",
"weight_block_size": [128, 128],
"modules_to_not_convert": ["model.visual", "lm_head"],
}
config.update(overrides)
return config
def test_h3_accepts_only_the_serialized_blockwise_checkpoint_contract() -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
assert config.weight_block_size == (128, 128)
assert config.get_supported_act_dtypes() == [torch.bfloat16]
with pytest.raises(ValueError, match=r"weight_block_size=\[128, 128\]"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(weight_block_size=[1, 128]))
with pytest.raises(ValueError, match="dynamic activation"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(activation_scheme="static"))
with pytest.raises(ValueError, match="vision stack"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(modules_to_not_convert=["lm_head"]))
with pytest.raises(ValueError, match="partially quantized language"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(modules_to_not_convert=["model.visual", "language_model.layers.3"]))
def test_serialized_fp8_allocates_checkpoint_weight_and_scale_without_requantization(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
layer = ColumnParallelLinear(
input_size=128,
output_size=256,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
)
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
assert layer.weight.dtype == torch.float8_e4m3fn
assert layer.weight.shape == (256, 128)
assert layer.weight_scale_inv.dtype == torch.float32
assert layer.weight_scale_inv.shape == (2, 1)
layer.weight.data.zero_()
layer.weight_scale_inv.data.fill_(0.25)
weight_pointer = layer.weight.data_ptr()
scale_pointer = layer.weight_scale_inv.data_ptr()
layer.quant_method.process_weights_after_loading(layer)
assert layer.weight.data_ptr() == weight_pointer
assert layer.weight_scale_inv.data_ptr() == scale_pointer
assert not hasattr(layer, "_fp8_weight")
def test_serialized_fp8_quantizes_only_language_linears(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
visual_linear = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.visual.blocks.0.attn.proj",
)
embedding = VocabParallelEmbedding(
num_embeddings=128,
embedding_dim=128,
org_num_embeddings=128,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.embed_tokens",
)
assert isinstance(visual_linear.quant_method, UnquantizedLinearMethod)
assert visual_linear.weight.dtype == torch.get_default_dtype()
assert isinstance(embedding.quant_method, UnquantizedEmbeddingMethod)
assert embedding.weight.dtype == torch.get_default_dtype()
def test_serialized_fp8_cpu_execution_fails_closed(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
layer = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.mlp.up_proj",
)
layer.weight.data.zero_()
layer.weight_scale_inv.data.fill_(1.0)
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
layer.quant_method.process_weights_after_loading(layer)
with pytest.raises(RuntimeError, match="requires CUDA"):
layer(torch.zeros(2, 128, dtype=torch.bfloat16))
def test_runtime_preflight_reports_capability_and_missing_dependencies(monkeypatch) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 0))
with pytest.raises(RuntimeError, match="sm100 or newer"):
config.validate_runtime(torch.device("cuda"))
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (10, 0))
def missing_quantizer() -> None:
raise RuntimeError("SGLang-compatible Triton quantizer is missing")
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", missing_quantizer)
with pytest.raises(RuntimeError, match="Triton quantizer is missing"):
config.validate_runtime(torch.device("cuda"))
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", lambda: None)
def missing_flashinfer():
raise RuntimeError("FlashInfer groupwise GEMM is missing")
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", missing_flashinfer)
with pytest.raises(RuntimeError, match="FlashInfer groupwise GEMM is missing"):
config.validate_runtime(torch.device("cuda"))
def test_loader_detects_and_capability_gates_checkpoint_metadata(tmp_path) -> None:
checkpoint_config = _checkpoint_quantization_config()
(tmp_path / "config.json").write_text(
json.dumps({"quantization_config": checkpoint_config}),
encoding="utf-8",
)
assert _read_text_encoder_checkpoint_quantization_config(str(tmp_path)) == checkpoint_config
model_config = MiniMaxH3Qwen3VLConfig()
quant_config = _configure_text_encoder_quantization(
model_config,
MiniMaxH3Qwen3VLConditioner,
str(tmp_path),
)
assert isinstance(quant_config, MiniMaxH3SerializedFP8Config)
assert model_config.quant_config is quant_config
unsupported_config = MiniMaxH3Qwen3VLConfig()
with pytest.raises(ValueError, match="does not support serialized 'fp8'"):
_configure_text_encoder_quantization(
unsupported_config,
TextEncoder,
str(tmp_path),
)
def test_loader_leaves_bf16_checkpoint_path_unchanged(tmp_path) -> None:
(tmp_path / "config.json").write_text(json.dumps({"architectures": ["Qwen3VLModel"]}), encoding="utf-8")
model_config = MiniMaxH3Qwen3VLConfig()
quant_config = _configure_text_encoder_quantization(
model_config,
MiniMaxH3Qwen3VLConditioner,
str(tmp_path),
)
assert quant_config is None
assert model_config.quant_config is None
def test_post_load_processing_visits_only_serialized_fp8_linears(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
quantized = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
)
plain = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
prefix="plain",
)
quantized.weight.data.zero_()
quantized.weight_scale_inv.data.fill_(1.0)
model = torch.nn.ModuleList([quantized, plain])
assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 1
assert quantized.weight.device.type == "cpu"
assert plain.weight.device.type == "cpu"
def test_flashinfer_groupwise_path_pins_output_dtype_and_trtllm_scale_layout(monkeypatch) -> None:
input_tensor = torch.zeros(2, 256, dtype=torch.bfloat16)
weight = torch.zeros(128, 256, dtype=torch.float8_e4m3fn)
weight_scale = torch.ones(1, 2, dtype=torch.float32)
quantized_input = torch.zeros_like(input_tensor, dtype=torch.float8_e4m3fn)
input_scale = torch.empty(2, 2, dtype=torch.float32).t()
input_scale.fill_(1.0)
receipt: dict[str, object] = {}
def fake_quantize(
value: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
assert value.data_ptr() == input_tensor.data_ptr()
assert value.shape == input_tensor.shape
assert group_size == 128
assert column_major_scales is True
return quantized_input, input_scale
def fake_gemm(
activation: torch.Tensor,
checkpoint_weight: torch.Tensor,
activation_scale: torch.Tensor,
checkpoint_scale: torch.Tensor,
*,
out_dtype: torch.dtype,
backend: str,
) -> torch.Tensor:
receipt.update(
activation=activation,
checkpoint_weight=checkpoint_weight,
activation_scale=activation_scale,
checkpoint_scale=checkpoint_scale,
out_dtype=out_dtype,
backend=backend,
)
return torch.zeros(activation.shape[0], checkpoint_weight.shape[0], dtype=out_dtype)
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_backend", lambda device: "trtllm")
monkeypatch.setattr(h3_fp8, "_sglang_per_token_group_quant_fp8", fake_quantize)
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", lambda: fake_gemm)
previous_default_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.float32)
try:
output = h3_fp8._flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor,
weight,
(128, 128),
weight_scale,
)
assert torch.get_default_dtype() == torch.float32
finally:
torch.set_default_dtype(previous_default_dtype)
assert output.dtype == torch.bfloat16
assert receipt["out_dtype"] == torch.bfloat16
assert receipt["backend"] == "trtllm"
assert receipt["activation"] is quantized_input
assert receipt["checkpoint_weight"] is weight
assert receipt["checkpoint_scale"] is weight_scale
assert receipt["activation_scale"] is input_scale
@@ -1,27 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
"""The Qwen3-VL stack is built only as far as MiniMax H3 reads.
H3 conditions on one intermediate hidden state. The layers above it were built,
weight-loaded and then discarded, which is 13.7 GB in bf16 and the difference
between fitting and not fitting on a 121 GB unified-memory device.
The dangerous part is not the truncation, it is getting the tuple index wrong.
`hidden_states` records each layer's *input*, so entry N is the output of layer
N-1, and the final entry comes from the norm that sits above the whole stack. A
truncated stack that still applies that norm puts a normalised tensor where the
raw one belongs: the length check in the conditioning stage still passes, and
conditioning silently changes. These tests pin the index, the content, and the
constant the two sides agree on.
"""
"""MiniMax-H3 Qwen3-VL layer truncation and slim-forward tests."""
from __future__ import annotations
import inspect
import os
import pytest
import torch
# Matches the other encoder tests: the module registry these build against wants
# a process group, and a single-rank one needs a rendezvous address.
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29513")
@@ -34,21 +20,15 @@ from fastvideo.pipelines.basic.minimax_h3.packing import MINIMAX_H3_TEXT_ENCODER
def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
"""A stack small enough to run on CPU but shaped like the real one.
Everything goes through the constructor so ``__post_init__`` validates the
small shape the same way it validates the real one.
"""
kwargs: dict = dict(
vocab_size=64,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=8,
output_hidden_state_index=5,
num_attention_heads=2,
num_key_value_heads=1,
head_dim=8,
# __post_init__ reads the sections out of rope_scaling, and they must
# cover exactly half of each head.
rope_scaling={
"mrope_interleaved": True,
"mrope_section": [2, 1, 1],
@@ -61,25 +41,27 @@ def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
def _small_config(**overrides) -> MiniMaxH3Qwen3VLConfig:
"""The outer config, which is what the modules take.
``ModelConfig.__getattr__`` forwards the architecture fields, so the modules
read ``prefix`` off this object and everything else off ``arch_config``.
"""
config = MiniMaxH3Qwen3VLConfig()
config.arch_config = _small_arch(**overrides)
return config
def test_default_matches_the_index_the_pipeline_reads() -> None:
"""The two sides cannot import each other, so pin them here instead.
config = MiniMaxH3Qwen3VLArchConfig()
`fastvideo/models/` must not import from `fastvideo/pipelines/`, so the tap
is written down twice. If they drift, conditioning reads a hidden state that
was never built and the run dies with an index error at generation time,
after a full model load.
"""
assert MiniMaxH3Qwen3VLArchConfig().num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
assert config.output_hidden_state_index == MINIMAX_H3_TEXT_ENCODER_LAYER
assert config.num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
def test_rejects_build_depth_that_cannot_reach_the_output() -> None:
for override in (0, 4):
with pytest.raises(ValueError, match="num_hidden_layers_override"):
_small_arch(num_hidden_layers_override=override)
def test_rejects_output_index_above_the_checkpoint_depth() -> None:
with pytest.raises(ValueError, match="output_hidden_state_index"):
_small_arch(output_hidden_state_index=9, num_hidden_layers_override=None)
def test_builds_only_up_to_the_override(distributed_setup) -> None:
@@ -87,7 +69,6 @@ def test_builds_only_up_to_the_override(distributed_setup) -> None:
assert model.num_layers == 5
assert len(model.layers) == 5
# The norm sits above the tap, so a truncated stack must not keep it.
assert model.norm is None
@@ -98,104 +79,94 @@ def test_override_none_keeps_the_full_stack(distributed_setup) -> None:
assert model.norm is not None
def test_nominal_and_built_depths_remain_distinct(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
assert conditioner.num_hidden_layers == 8
assert conditioner.num_built_hidden_layers == 5
def test_override_above_the_stack_does_not_over_build(distributed_setup) -> None:
# num_hidden_layers comes from the checkpoint's config.json via
# update_model_arch, so a smaller variant must clamp rather than ask for
# layers that do not exist.
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=99))
assert model.num_layers == 8
assert model.norm is not None
def test_override_equal_to_the_stack_keeps_the_norm(distributed_setup) -> None:
"""The exact boundary of the clamp: a stack cut at its own depth is full.
A checkpoint with exactly ``override`` layers taps its final layer, whose
tuple entry sits after the norm in the full model, so the norm must stay
and nothing may be filtered from the checkpoint.
"""
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=8))
assert model.num_layers == 8
assert model.norm is not None
def test_non_positive_override_is_rejected() -> None:
"""A non-positive override would build no decoder layers at all.
Worse, a negative one makes ``num_layers`` disagree with the built stack
and the surplus-key filter would then drop every layer key, so the
conditioner would load "successfully" with no transformer. Reject it at
config construction, and again when update_model_arch re-validates.
"""
for override in (0, -1):
with pytest.raises(ValueError, match="num_hidden_layers_override"):
_small_arch(num_hidden_layers_override=override)
config = _small_config()
with pytest.raises(ValueError, match="num_hidden_layers_override"):
config.update_model_arch({"num_hidden_layers_override": 0})
def test_tapped_hidden_state_is_unchanged_by_truncation(distributed_setup) -> None:
"""The whole point: entry `tap` must be bit-identical either way."""
"""The slim model returns the raw output at the selected layer."""
tap = 5
full = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
cut = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=tap))
# These modules allocate uninitialised storage and expect a checkpoint, so
# give them finite weights before running anything through them.
torch.manual_seed(0)
for parameter in full.parameters():
parameter.data.normal_(std=0.02)
# Then make the shared prefix identical, which is the only part the tapped
# hidden state depends on.
for (_, a), (_, b) in zip(full.layers[:tap].named_parameters(),
cut.layers[:tap].named_parameters(),
strict=True):
b.data.copy_(a.data)
torch.manual_seed(1)
inputs_embeds = torch.randn(1, 6, 16)
# mRoPE indexes three axes (t, h, w); text tokens share the same position on
# all three.
position_ids = torch.arange(6).view(1, 1, 6).expand(3, 1, 6)
with torch.no_grad():
full_out = full(inputs_embeds, position_ids, None, True, None, None)
cut_out = cut(inputs_embeds, position_ids, None, True, None, None)
expected = inputs_embeds
position_embeddings = full.rotary_emb(inputs_embeds, position_ids)
for layer in full.layers[:tap]:
expected = layer(expected, position_embeddings, None)
full_out = full(inputs_embeds, position_ids, None, None, None)
cut_out = cut(inputs_embeds, position_ids, None, None, None)
assert torch.equal(full_out.hidden_states[tap], cut_out.hidden_states[tap])
# And the truncated model must not offer states it never computed.
assert len(cut_out.hidden_states) == tap + 1
# The whole shared prefix must match, not just the tap: this is the same
# comparison the production-loader parity gate runs against the official
# model, and it is what catches a truncated stack that still applied the
# final norm to its last entry.
for index, (cut_state, full_state) in enumerate(zip(cut_out.hidden_states, full_out.hidden_states,
strict=False)):
assert torch.equal(cut_state, full_state), f"hidden state {index} changed under truncation"
assert torch.equal(expected, full_out)
assert torch.equal(expected, cut_out)
def test_conditioning_stage_adapts_slim_sequence_output() -> None:
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage
class FakeConditioner:
dtype = torch.float32
def __call__(self, input_ids: torch.Tensor, **kwargs) -> torch.Tensor:
assert input_ids.ndim == 1
assert not kwargs
return torch.ones(input_ids.shape[0], 4)
stage = MiniMaxH3ConditioningStage(conditioner=FakeConditioner(), tokenizer=None, processor=None, ref2va=False)
embeddings, tags = stage._encode_tokens([1, 2, 3], [0, 0, 0], torch.device("cpu"))
assert embeddings.shape == (1, 3, 4)
assert tags.shape == (3, )
def test_conditioner_exposes_only_the_slim_forward_contract() -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
assert tuple(inspect.signature(MiniMaxH3Qwen3VLConditioner.forward).parameters) == (
"self",
"input_ids",
"pixel_values",
"image_grid_thw",
"pixel_values_videos",
"video_grid_thw",
)
def test_truncated_model_drops_the_surplus_checkpoint_keys(distributed_setup) -> None:
"""The unexpected-key check is strict on purpose, so the surplus keys have
to be filtered rather than the check relaxed."""
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
assert conditioner._is_above_the_tap("language_model.layers.5.mlp.gate_proj.weight")
assert conditioner._is_above_the_tap("language_model.layers.7.self_attn.q_proj.weight")
assert conditioner._is_above_the_tap("language_model.norm.weight")
# Kept: layers we built, the embeddings, and the vision tower.
assert not conditioner._is_above_the_tap("language_model.layers.4.mlp.gate_proj.weight")
assert not conditioner._is_above_the_tap("language_model.embed_tokens.weight")
assert not conditioner._is_above_the_tap("visual.blocks.0.attn.qkv.weight")
# The filter only drops indexes the full stack would have built. A key at
# or above the checkpoint's own num_hidden_layers is corrupt, and it must
# keep raising as unexpected exactly as it does without truncation.
assert not conditioner._is_above_the_tap("language_model.layers.8.mlp.gate_proj.weight")
with pytest.raises(ValueError, match="Unexpected"):
conditioner.load_weights([("model.language_model.layers.8.mlp.gate_proj.weight", torch.zeros(1))])
assert conditioner._is_omitted_checkpoint_key("language_model.layers.5.mlp.gate_proj.weight")
assert conditioner._is_omitted_checkpoint_key("language_model.layers.7.self_attn.q_proj.weight")
assert conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.4.mlp.gate_proj.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.embed_tokens.weight")
assert not conditioner._is_omitted_checkpoint_key("visual.blocks.0.attn.qkv.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.8.mlp.gate_proj.weight")
def test_full_stack_filters_nothing(distributed_setup) -> None:
@@ -203,5 +174,14 @@ def test_full_stack_filters_nothing(distributed_setup) -> None:
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=None))
assert not conditioner._is_above_the_tap("language_model.layers.7.mlp.gate_proj.weight")
assert not conditioner._is_above_the_tap("language_model.norm.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.7.mlp.gate_proj.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
def test_corrupt_layer_above_checkpoint_depth_remains_unexpected(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
with pytest.raises(ValueError, match="Unexpected"):
conditioner.load_weights([("language_model.layers.8.mlp.gate_proj.weight", torch.empty(1))])
@@ -2,8 +2,10 @@ import os
from types import SimpleNamespace
import warnings
import numpy as np
import pytest
import torch
from einops import rearrange
import fastvideo.entrypoints.video_generator as video_generator_module
from fastvideo.api import (
@@ -271,6 +273,70 @@ def test_generate_single_video_return_frames_still_materializes_output(tmp_path)
assert result["video_path"] is None
def test_generate_single_video_frames_match_legacy_cpu_loop(tmp_path):
"""The on-device quantize path (#1362) must reproduce the legacy
per-frame CPU loop (make_grid -> permute -> *255 -> uint8) bit-exactly
for in-range fp32 pixels: same uint8 dtype, same HWC grid layout with
nrow=6 (batch>1), odd frame count. CPU-only: on CUDA the float->uint8
cast may differ by <=1 LSB, but on CPU both orderings run identical
fp32 ops, so exact equality is required."""
torch.manual_seed(0)
output = torch.rand((2, 3, 3, 16, 16), dtype=torch.float32)
output_batch = _single_video_output_batch(output)
fastvideo_args = _single_video_args()
generator = _single_video_generator(output_batch, fastvideo_args)
sampling_param = _small_sampling_param(save_video=False, return_frames=True)
sampling_param.num_frames = 3
sampling_param.num_videos_per_prompt = 2
result = generator._generate_single_video(
prompt="grid parity",
sampling_param=sampling_param,
fastvideo_args=fastvideo_args,
output_path=str(tmp_path / "unused.mp4"),
)
legacy_frames = []
for x in rearrange(output, "b c t h w -> t b c h w"):
grid = video_generator_module.torchvision.utils.make_grid(x, nrow=6)
grid = grid.permute(1, 2, 0).squeeze(-1)
legacy_frames.append((grid * 255).to(torch.uint8).contiguous().cpu().numpy())
torch.testing.assert_close(result["samples"], output)
assert len(result["frames"]) == 3
for got, want in zip(result["frames"], legacy_frames, strict=True):
assert got.dtype == np.uint8
assert got.shape == want.shape
np.testing.assert_array_equal(got, want)
def test_generate_single_video_frames_clamp_out_of_range_pixels(tmp_path):
"""VAE output slightly outside [0, 1] must saturate at 0/255 in the
uint8 frames. The pre-#1362 unclamped cast wrapped mod 256 (e.g.
1.5 -> 126). CPU-only."""
output = torch.full((1, 3, 2, 16, 16), 1.5, dtype=torch.float32)
output[:, :, 1] = -0.5
output_batch = _single_video_output_batch(output)
fastvideo_args = _single_video_args()
generator = _single_video_generator(output_batch, fastvideo_args)
result = generator._generate_single_video(
prompt="clamp",
sampling_param=_small_sampling_param(save_video=False, return_frames=True),
fastvideo_args=fastvideo_args,
output_path=str(tmp_path / "unused.mp4"),
)
frames = result["frames"]
assert len(frames) == 2
# make_grid passes a single image through without grid padding, so
# every pixel comes from the (clamped) output tensor.
assert frames[0].dtype == np.uint8
assert frames[0].shape == (16, 16, 3)
assert (frames[0] == 255).all()
assert (frames[1] == 0).all()
def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_path):
output = torch.ones((1, 3, 2, 16, 16), dtype=torch.float32) * 0.5
output_batch = _single_video_output_batch(output)
@@ -305,6 +371,37 @@ def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_p
}
def test_generate_single_video_save_only_reports_refined_output_size(monkeypatch, tmp_path):
"""`GenerationResult.size` must describe the decoded media even when the
fp32 `samples` mirror is skipped (`return_frames=False`, the CLI save
flow). Refiner pipelines can change the final pixel geometry, so the size
has to come from `output_batch.output`, not the base request. CPU-only."""
# Refiner-style output: request asks for 2 frames of 16x16, pipeline
# produces 5 frames of 32x48.
output = torch.full((1, 3, 5, 32, 48), 0.5, dtype=torch.float32)
output_batch = _single_video_output_batch(output)
fastvideo_args = _single_video_args()
generator = _single_video_generator(output_batch, fastvideo_args)
saved = {}
def fake_mimsave(path, frames, *, fps, format):
saved["frame_count"] = len(frames)
monkeypatch.setattr(video_generator_module.imageio, "mimsave", fake_mimsave)
result = generator._generate_single_video(
prompt="refined save",
sampling_param=_small_sampling_param(save_video=True, return_frames=False),
fastvideo_args=fastvideo_args,
output_path=str(tmp_path / "refined.mp4"),
)
assert result["samples"] is None
assert result["frames"] is None
assert result["size"] == (32, 48, 5)
assert saved["frame_count"] == 5
def test_generate_single_video_audio_only_metadata_returns_audio_without_frames(tmp_path):
audio = torch.zeros((16, ), dtype=torch.float32)
output_batch = _single_video_output_batch(

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