Compare commits

...
Author SHA1 Message Date
Will LinandClaude Fable 5 9ba740e125 [docs] inference optimizations: document the regional fullgraph compile knob
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:55:03 +00:00
Will LinandClaude Fable 5 34346779d1 [perf] inference: regional fullgraph torch.compile of the DiT blocks (port of #1718 to the inference path)
Extend the training-side regional-compile port (internal 863e87342, GO
verdict vsa_gate/compile_ab/VERDICT.md) to inference so every H3 run can
compile the 52 transformer blocks. Opt-in and off by default:

- New FastVideoArgs.inference_torch_compile (CLI --inference-torch-compile),
  env FASTVIDEO_INFERENCE_TORCH_COMPILE=1 folded in __post_init__ (the
  attention_backend parse-once pattern), reachable through
  PipelineSelection.experimental {"inference_torch_compile": true} exactly
  like the VSA_sparsity / VSA_tile_size knobs.
- maybe_load_fsdp_model applies the compile right after the transformer
  loads: per-_compile_conditions block, fullgraph=True + inductor
  options.emulate_precision_casts injected, no user kwargs needed (the
  compile-A/B verdict recipe). prepare_for_compile still runs first, so the
  H3 fusion-inertness warning of #1735 is preserved.
- _regional_compile_unsupported_reason ported from the training loader:
  VSA / VSA-H3 backends and the FASTVIDEO_DISABLE_ATTENTION_COMPILE=1
  escape hatch degrade the transformer to eager with one warning; FLASH_ATTN
  on flash-attn 3 is rejected with an actionable message. FA4/FA2/SDPA
  compile through their existing custom-op boundaries (forward-only is
  enough at inference; the training port's backward custom ops are not
  needed under no_grad).
- FASTVIDEO_DISABLE_ATTENTION_COMPILE default flipped to 0 (attention is
  traced by default), matching the training port and upstream #1718 —
  eager runs are unaffected (torch.compiler.disable is inert outside
  dynamo).
- enable_torch_compile + inference_torch_compile together skip the
  pipeline-level DiT compile (the loader already owns those forwards).
- examples: --inference-torch-compile on basic_minimax_h3_t2v.py (dense
  target) and basic_fasth3.py (exercises the VSA degrade path).
- CPU-safe contract tests for the guard and the kwargs injection.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:46:26 +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
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
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
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
H1yori233 74b409d7cf [test]: add H3 VAE parity and memory benchmark 2026-08-19 02:45:23 -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
44 changed files with 4361 additions and 431 deletions
+25
View File
@@ -337,6 +337,31 @@ Only DiT submodules that declare `_compile_conditions` are compiled
(most shipped models). The text encoder and VAE are not compiled by this
flag.
### Regional fullgraph compile (experimental)
`inference_torch_compile` is a stricter, kwargs-free variant that ports the
training-side regional compile of
[#1718](https://github.com/hao-ai-lab/FastVideo/pull/1718) to inference: the
loader wraps each `_compile_conditions` block in
`torch.compile(fullgraph=True)` with inductor
`options={"emulate_precision_casts": True}` right after the transformer
loads. Attention backends that cannot be traced end-to-end (VSA,
FLASH_ATTN on flash-attn 3, or the `FASTVIDEO_DISABLE_ATTENTION_COMPILE=1`
escape hatch) degrade the transformer to eager with one warning instead of
failing mid-denoise.
```python
generator = VideoGenerator.from_pretrained(
"MiniMaxAI/MiniMax-H3",
inference_torch_compile=True, # or FASTVIDEO_INFERENCE_TORCH_COMPILE=1
)
```
Do not combine it with `torch_compile_kwargs['mode']` (the loader injects
inductor options, and torch.compile forbids mode+options); it is
independent of `enable_torch_compile`, and when both are set the regional
compile wins for the DiT.
### What to expect
| Config | Effect |
+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!
+188
View File
@@ -0,0 +1,188 @@
# 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("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load. NOTE: this "
"script always runs the VSA-H3 backend, which is not fullgraph-traceable — the "
"loader logs one warning and keeps the transformer eager. The flag is exposed "
"here to exercise exactly that guard")
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
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
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()
@@ -15,6 +15,7 @@ from fastvideo.api import (
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
@@ -41,6 +42,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load (the #1718 "
"training-port semantics: no kwargs; fullgraph + emulate_precision_casts injected). "
"First generation pays the inductor JIT (~1-2 min); use --repeats >= 2 and time "
"the last repeat. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent")
parser.add_argument("--repeats",
type=int,
default=1,
@@ -54,9 +61,16 @@ def main() -> None:
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
# Boot-time run configuration folded into FastVideoArgs (the same
# experimental-dict route basic_fasth3.py uses for the VSA knobs).
experimental: dict[str, object] = {}
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
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,
+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,14 +433,66 @@ 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, seq_len, heads, dim = out.shape
@@ -325,7 +501,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
# 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_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)]
+4 -4
View File
@@ -18,13 +18,13 @@ from fastvideo.layers.rotary_embedding import _apply_rotary_emb
def _attention_compile_disabled() -> bool:
"""Whether to keep attention ``forward`` out of the torch.compile graph.
Defaults to ``True`` (the historical behavior: attention runs eager via
``torch.compiler.disable``). Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=0``
to let attention be traced/compiled into the surrounding graph.
Attention backends expose traceable custom-op boundaries, so compilation
is enabled by default. Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=1`` for
an explicit eager escape hatch when debugging a backend.
"""
val = os.environ.get("FASTVIDEO_DISABLE_ATTENTION_COMPILE")
if val is None:
return True
return False
return val.strip().lower() not in ("0", "false", "no", "off", "")
@@ -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))
+25
View File
@@ -21,12 +21,15 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_INFERENCE_TORCH_COMPILE: bool = False
FASTVIDEO_MINIMAX_H3_FUSIONS: str = ""
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 +220,32 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# If set (=1), enable regional (per-transformer-block) fullgraph
# torch.compile for the DiT at inference — the inference-side counterpart
# of the training regional-compile port of hao-ai-lab/FastVideo#1718.
# Equivalent to FastVideoArgs.inference_torch_compile=True (e.g. via
# PipelineSelection.experimental={"inference_torch_compile": True}). VSA
# and other non-fullgraph-traceable attention backends degrade to eager
# with one warning; see _regional_compile_unsupported_reason in
# fastvideo/models/loader/fsdp_load.py.
"FASTVIDEO_INFERENCE_TORCH_COMPILE":
lambda: os.getenv("FASTVIDEO_INFERENCE_TORCH_COMPILE", "0") != "0",
# 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":
+35
View File
@@ -164,11 +164,24 @@ class FastVideoArgs:
torch_compile_kwargs_text_encoder: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_vae: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_audio_vae: dict[str, Any] = field(default_factory=dict)
# Regional (per-transformer-block) fullgraph torch.compile of the DiT at
# inference — the inference-side counterpart of the training regional
# compile ported from hao-ai-lab/FastVideo#1718. Applied by the loader
# right after the transformer loads, with fullgraph=True and inductor
# options {emulate_precision_casts: True} injected (no user kwargs
# needed). Attention backends that cannot be fullgraph-traced (VSA,
# FLASH_ATTN on flash-attn 3) degrade the transformer to eager with one
# warning. Opt-in via FASTVIDEO_INFERENCE_TORCH_COMPILE=1 (folded in
# __post_init__) or PipelineSelection.experimental
# {"inference_torch_compile": true}. Distinct from ``enable_torch_compile``,
# which keeps the pipeline-level compile semantics.
inference_torch_compile: bool = False
disable_autocast: bool = False
# 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
@@ -271,6 +284,13 @@ class FastVideoArgs:
self._apply_ltx2_vae_overrides()
self._resolve_refine_args()
self._apply_transformer_quant()
if not self.inference_torch_compile:
# Parse-once adapter (same pattern as attention_backend below): the
# environment variable is an input read once here, so the loader
# only ever consults the typed field.
import fastvideo.envs as envs
if envs.FASTVIDEO_INFERENCE_TORCH_COMPILE:
self.inference_torch_compile = True
if self.attention_backend is not None:
# Fail fast on typos instead of silently auto-selecting later.
from fastvideo.attention.selector import coerce_attn_backend
@@ -591,6 +611,15 @@ class FastVideoArgs:
help=
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
)
parser.add_argument(
"--inference-torch-compile",
action=StoreBoolean,
default=FastVideoArgs.inference_torch_compile,
help="Regional fullgraph torch.compile of each DiT transformer block at inference "
"(port of the #1718 training-side regional compile). The loader injects fullgraph=True "
"and inductor options {emulate_precision_casts: true}; non-traceable attention backends "
"(VSA) degrade to eager with one warning. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent.",
)
parser.add_argument(
"--dit-cpu-offload",
@@ -644,6 +673,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(
+136 -26
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 "
@@ -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",
]
+62 -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={
@@ -1080,6 +1126,7 @@ class TransformerLoader(ComponentLoader):
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
inference_regional_compile=fastvideo_args.inference_torch_compile,
)
total_params = sum(p.numel() for p in model.parameters())
+117
View File
@@ -136,6 +136,7 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
inference_regional_compile: bool = False,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -235,9 +236,125 @@ def maybe_load_fsdp_model(
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
elif inference_regional_compile and not training_mode:
# Inference-side counterpart of the #1718 training regional compile:
# per-block fullgraph compile right after the transformer loads, no
# user kwargs needed (fullgraph + emulate_precision_casts injected).
unsupported = _regional_compile_unsupported_reason(init_params)
if unsupported is not None:
logger.warning(
"inference_torch_compile requested but disabled: %s. "
"Inference continues in eager mode.", unsupported)
else:
prepare_for_compile = getattr(model, "prepare_for_compile", None)
if callable(prepare_for_compile):
logger.info("Running prepare_for_compile for %s", type(model).__name__)
prepare_for_compile()
_compile_model_regions(model, torch_compile_kwargs or {})
return model
def _regional_compile_unsupported_reason(init_params: dict[str, Any]) -> str | None:
"""Return why regional fullgraph compile cannot run, or None if it can.
FA3's grad-enabled attention path deliberately routes to the raw
autograd.Function at the cost of a dynamo graph break (see
flash_attn_default.py) — under regional ``fullgraph=True`` that break is
a hard RuntimeError at the first compiled forward. FA2 and FA4 route
through traceable custom ops and are compile-safe.
The VSA backends (Triton block-sparse kernels behind sequence-parallel
all-to-alls plus a host-synced metadata guard) are likewise not
fullgraph-traceable; a VSA-backed transformer (e.g. the FastH3 student)
falls back to eager while compile-safe dense loads still compile.
"""
try:
from fastvideo.attention.layer import _attention_compile_disabled
except Exception: # pragma: no cover - attention stack not importable
pass
else:
if _attention_compile_disabled():
# The escape hatch wraps attention forwards in
# torch.compiler.disable, which is a hard dynamo error inside a
# fullgraph region ("Skip inlining `torch.compiler.disable()`d
# function"). Degrade to eager instead, matching the hatch's
# debugging intent.
return ("FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 keeps attention "
"forwards out of compiled graphs via torch.compiler."
"disable, which fullgraph regional compile cannot trace; "
"this model stays eager")
config = init_params.get("config")
resolved = getattr(config, "_resolved_attention_backend", None)
resolved_name = getattr(resolved, "name", "")
if resolved_name in ("VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"):
return (f"attention backend resolved to {resolved_name}, whose Triton "
"kernels, sequence-parallel collectives, and sync metadata "
"guard graph-break (incompatible with fullgraph regional "
"compile); this model stays eager")
if resolved is None or resolved_name != "FLASH_ATTN":
return None
try:
from fastvideo.attention.utils.flash_attn_default import fa_version
except Exception: # pragma: no cover - flash-attn stack not importable
return None
if fa_version == "3":
return ("attention backend resolved to FLASH_ATTN with flash-attn 3, "
"whose grad-enabled path graph-breaks (incompatible with "
"fullgraph regional compile); use FA2, FA4 (FASTVIDEO_FA4=1), "
"or TORCH_SDPA for compiled runs")
return None
def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> int:
"""Compile repeated mathematical regions of a loaded model.
Only the selected module ``forward`` is replaced. This keeps activation
checkpoint wrappers structurally transparent while any module-level hooks
(FSDP pre/post, layerwise offload) execute outside the compiled region.
"""
compile_conditions = getattr(model, "_compile_conditions", None)
if not compile_conditions:
raise ValueError(f"{type(model).__name__} does not declare _compile_conditions")
if compile_kwargs.get("fullgraph", True) is not True:
raise ValueError("Regional compile requires fullgraph=True")
if "mode" in compile_kwargs:
# torch.compile forbids passing both `mode` and `options`, and
# regional compile always injects options (emulate_precision_casts)
# for bf16 numerics parity. Fail here with an actionable message
# instead of letting torch raise a mode/options conflict about an
# `options` key the user never wrote.
raise ValueError("Regional compile sets inductor options "
"(emulate_precision_casts) and cannot be combined "
"with torch_compile_kwargs['mode']. Remove 'mode' or "
"express its effect via torch_compile_kwargs['options'].")
kwargs = {**compile_kwargs, "fullgraph": True}
options = {"emulate_precision_casts": True}
options.update(kwargs.get("options") or {})
kwargs["options"] = options
compiled_count = 0
for name, submodule in list(model.named_modules()):
if not name:
continue
if any(condition(name, submodule) for condition in compile_conditions):
# Activation checkpoint wrappers are control-flow boundaries, not
# mathematical regions. Keep their saved-tensor/recompute logic
# eager and compile only the repeated block they own.
compile_target = getattr(submodule, "_checkpoint_wrapped_module", submodule)
compile_target.forward = torch.compile(compile_target.forward, **kwargs)
compiled_count += 1
if compiled_count == 0:
raise ValueError(f"No submodules in {type(model).__name__} matched _compile_conditions")
logger.info(
"Enabled regional torch.compile for %d submodules in %s with kwargs=%s",
compiled_count,
type(model).__name__,
kwargs,
)
return compiled_count
def shard_model(
model,
*,
@@ -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
+128 -44
View File
@@ -15,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:
@@ -292,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
@@ -303,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))
@@ -329,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)
@@ -434,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,
@@ -483,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."""
@@ -490,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__()
@@ -655,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 = []
@@ -677,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))
@@ -700,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
@@ -809,21 +887,27 @@ class AutoencoderKLMiniMaxH3(nn.Module):
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])
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
next_overlap = None
if self.config.token_drop > 0:
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
if overlap is not None:
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
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)
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
chunk = chunk[:, :, :num_frames]
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
if num_frames > 0:
yield chunk[:, :, :num_frames]
output_frame_start += num_frames
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]
@@ -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")
@@ -11,6 +11,7 @@ from fastvideo.distributed import get_local_torch_device, get_world_group, model
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE
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,
@@ -55,6 +56,7 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 video latents into normalized CPU pixels."""
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
# Distributed executors consume rank 0's ForwardBatch. Keep a
# verifier-compatible placeholder on other ranks and avoid
@@ -88,8 +90,12 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
dtype=torch.float32,
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
)
# The published decode recipe uses FP16 autocast over FP32 weights.
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
# 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"),
):
self.vae.decode_to_pixels(latents, output)
batch.output = output
return batch
@@ -121,6 +127,7 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 audio latents into a stereo CPU waveform."""
if model_parallel_is_initialized() and not get_world_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
@@ -144,7 +151,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
@@ -203,6 +203,14 @@ class ComposedPipelineBase(ABC):
vae_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_vae or global_compile_kwargs)
audio_vae_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_audio_vae or global_compile_kwargs)
if compile_transformer and self.fastvideo_args.inference_torch_compile:
# The loader already applied the regional fullgraph
# compile to the DiT blocks (inference_torch_compile);
# wrapping the same forwards again here would stack
# compiled callables.
logger.info("inference_torch_compile already compiled the DiT regions in the "
"loader; skipping the pipeline-level DiT compile")
compile_transformer = False
if compile_transformer:
self._maybe_compile_pipeline_module(
module_name="transformer",
+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."""
@@ -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))])
@@ -0,0 +1,117 @@
# SPDX-License-Identifier: Apache-2.0
"""Contract tests for the inference-side regional torch.compile port.
The loader applies a per-transformer-block fullgraph compile after the
transformer loads (``FastVideoArgs.inference_torch_compile``, env
``FASTVIDEO_INFERENCE_TORCH_COMPILE=1``). These tests pin the two pieces that
must not drift from the #1718 training-port semantics:
- ``_regional_compile_unsupported_reason``: VSA backends (and the attention
eager escape hatch) degrade to eager with a reason instead of hard-failing
fullgraph capture at the first denoising forward.
- ``_compile_model_regions``: exactly the ``_compile_conditions`` blocks are
compiled, fullgraph=True plus inductor ``emulate_precision_casts`` are
injected, and ``mode`` kwargs are rejected (torch.compile forbids
mode+options).
CPU-safe: torch.compile is monkeypatched, no CUDA needed.
"""
from types import SimpleNamespace
import pytest
import torch
from torch import nn
from fastvideo.models.loader import fsdp_load
from fastvideo.models.loader.fsdp_load import (
_compile_model_regions,
_regional_compile_unsupported_reason,
)
def _init_params_for(backend_name: str | None) -> dict:
resolved = None if backend_name is None else SimpleNamespace(name=backend_name)
return {"config": SimpleNamespace(_resolved_attention_backend=resolved)}
@pytest.mark.parametrize("backend_name", ["VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"])
def test_vsa_backends_degrade_to_eager(backend_name, monkeypatch) -> None:
monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False)
reason = _regional_compile_unsupported_reason(_init_params_for(backend_name))
assert reason is not None
assert backend_name in reason
assert "eager" in reason
@pytest.mark.parametrize("backend_name", [None, "TORCH_SDPA"])
def test_dense_backends_allow_compile(backend_name, monkeypatch) -> None:
monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False)
assert _regional_compile_unsupported_reason(_init_params_for(backend_name)) is None
def test_attention_compile_escape_hatch_degrades_to_eager(monkeypatch) -> None:
monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1")
reason = _regional_compile_unsupported_reason(_init_params_for("TORCH_SDPA"))
assert reason is not None
assert "FASTVIDEO_DISABLE_ATTENTION_COMPILE" in reason
class _Block(nn.Module):
def __init__(self) -> None:
super().__init__()
self.linear = nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x)
class _Toy(nn.Module):
_compile_conditions = [lambda name, module: name.startswith("blocks.") and name.count(".") == 1]
def __init__(self) -> None:
super().__init__()
self.blocks = nn.ModuleList([_Block() for _ in range(3)])
self.proj_out = nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for block in self.blocks:
x = block(x)
return self.proj_out(x)
def test_compile_model_regions_injects_fullgraph_and_precision_casts(monkeypatch) -> None:
captured: list[dict] = []
def _fake_compile(fn, **kwargs):
captured.append(kwargs)
return fn
monkeypatch.setattr(fsdp_load.torch, "compile", _fake_compile)
model = _Toy()
count = _compile_model_regions(model, {})
# The three repeated blocks compile; proj_out and the root stay eager.
assert count == 3
assert len(captured) == 3
for kwargs in captured:
assert kwargs["fullgraph"] is True
assert kwargs["options"] == {"emulate_precision_casts": True}
def test_compile_model_regions_rejects_mode_kwargs() -> None:
with pytest.raises(ValueError, match="mode"):
_compile_model_regions(_Toy(), {"mode": "reduce-overhead"})
def test_compile_model_regions_requires_conditions_and_matches(monkeypatch) -> None:
monkeypatch.setattr(fsdp_load.torch, "compile", lambda fn, **kwargs: fn)
plain = nn.Linear(4, 4)
with pytest.raises(ValueError, match="_compile_conditions"):
_compile_model_regions(plain, {})
class _NoMatch(_Toy):
_compile_conditions = [lambda name, module: False]
with pytest.raises(ValueError, match="matched"):
_compile_model_regions(_NoMatch(), {})
@@ -0,0 +1,231 @@
# SPDX-License-Identifier: Apache-2.0
"""Focused routing and FA4 integration checks for MiniMax-H3 fusions."""
from __future__ import annotations
import pytest
import torch
from fastvideo.platforms import AttentionBackendEnum
@pytest.mark.parametrize(
("raw", "expected"),
[
("", frozenset()),
("0", frozenset()),
("none", frozenset()),
("all", frozenset({"modulate", "qknorm_rope", "swiglu"})),
("1", frozenset({"modulate", "qknorm_rope", "swiglu"})),
("swiglu, modulate", frozenset({"swiglu", "modulate"})),
],
)
def test_minimax_h3_fusion_selector(raw: str, expected: frozenset[str]) -> None:
from fastvideo.models.dits.minimax_h3 import _enabled_minimax_h3_fusions
assert _enabled_minimax_h3_fusions(raw) == expected
def test_minimax_h3_fusion_selector_rejects_unknown_name() -> None:
from fastvideo.models.dits.minimax_h3 import _enabled_minimax_h3_fusions
with pytest.raises(ValueError, match="Unknown MiniMax H3 fusion"):
_enabled_minimax_h3_fusions("swiglu,unknown")
def test_swiglu_fusion_stays_on_eager_path_with_grad(monkeypatch: pytest.MonkeyPatch) -> None:
import fastvideo.models.dits.minimax_h3 as h3
def unexpected_kernel(_: torch.Tensor) -> torch.Tensor:
raise AssertionError("inference-only fusion ran with grad enabled")
monkeypatch.setattr(h3, "minimax_h3_swiglu", unexpected_kernel)
layer = h3.MiniMaxH3FeedForward(8, 16, fuse_swiglu=True)
inputs = torch.randn(2, 3, 8, requires_grad=True)
layer(inputs).sum().backward()
assert inputs.grad is not None
def test_all_minimax_h3_fusions_match_one_eager_block_under_fa4(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Exercise the real block wiring without loading any H3 checkpoint."""
if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
pytest.skip("BF16 CUDA is required")
pytest.importorskip("triton")
flash_attn = pytest.importorskip("flash_attn")
if "fa4" not in getattr(flash_attn, "__version__", "").lower():
pytest.skip("the focused integration test requires the FA4 environment")
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
monkeypatch.setenv("FASTVIDEO_FA4", "1")
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
monkeypatch.setenv("MASTER_PORT", "29573")
monkeypatch.setenv("RANK", "0")
monkeypatch.setenv("WORLD_SIZE", "1")
monkeypatch.setenv("LOCAL_RANK", "0")
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
from fastvideo.forward_context import set_forward_context
from fastvideo.models.dits.minimax_h3 import MiniMaxH3RotaryPosEmbed, MiniMaxH3TransformerBlock
maybe_init_distributed_environment_and_model_parallel(1, 1)
try:
kwargs = dict(
hidden_size=128,
num_attention_heads=1,
attention_head_dim=128,
ffn_dim=256,
time_embed_dim=64,
norm_eps=1e-5,
qk_norm_eps=1e-5,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, ),
quant_config=None,
prefix="minimax_h3.test_block",
)
previous_default_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.bfloat16)
try:
eager = MiniMaxH3TransformerBlock(**kwargs)
fused = MiniMaxH3TransformerBlock(
**kwargs,
fuse_modulate=True,
fuse_qknorm_rope=True,
fuse_swiglu=True,
)
finally:
torch.set_default_dtype(previous_default_dtype)
with torch.no_grad():
for name, parameter in eager.named_parameters():
if "norm" in name and name.endswith("weight"):
parameter.fill_(1.0)
elif parameter.ndim > 1:
torch.nn.init.normal_(parameter, mean=0.0, std=0.02)
else:
parameter.zero_()
fused.load_state_dict(eager.state_dict(), strict=True)
device = torch.device("cuda")
eager = eager.to(device=device, dtype=torch.bfloat16).eval()
fused = fused.to(device=device, dtype=torch.bfloat16).eval()
generator = torch.Generator(device=device).manual_seed(2026)
hidden_states = torch.randn(2, 12, 128, generator=generator, device=device, dtype=torch.bfloat16)
temb = torch.randn(2, 64, generator=generator, device=device, dtype=torch.bfloat16)
adaln_indices = torch.arange(12, device=device, dtype=torch.long).remainder(6)
position_ids = torch.zeros(12, 3, device=device, dtype=torch.float32)
position_ids[:, 0] = torch.arange(12, device=device)
rotary_emb = MiniMaxH3RotaryPosEmbed(rope_freq_dim=16, rope_theta=10000.0).to(device)(position_ids)
inputs = dict(
hidden_states=hidden_states,
temb=temb,
adaln_indices=adaln_indices,
rotary_emb=tuple(value.to(torch.bfloat16) for value in rotary_emb),
original_seq_len=12,
)
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
eager_output = eager(**inputs)
fused_output = fused(**inputs)
# Sol-Engine keeps fused intermediates in FP32 registers until their
# final BF16 stores, so the opt-in path is close but not bit-identical.
torch.testing.assert_close(fused_output, eager_output, atol=3e-2, rtol=3e-2)
finally:
cleanup_dist_env_and_memory()
def test_minimax_h3_fusions_engage_on_cuda_inference(monkeypatch: pytest.MonkeyPatch) -> None:
"""Pin the positive side of the routing guard.
The parity test above still passes if ``_can_run_minimax_h3_fusion``
silently degrades to always-False (both blocks then run the identical
eager path), so count the fused-kernel calls: one CUDA inference forward
through a fully fused block must hit ``fused_rmsnorm_modulate`` once,
``fused_residual_gate_rmsnorm_modulate`` once, ``fused_qknorm_rope``
twice (q and k), and ``minimax_h3_swiglu`` once -- and a grad-enabled
forward must leave every counter unchanged.
"""
if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
pytest.skip("BF16 CUDA is required")
pytest.importorskip("triton")
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
monkeypatch.setenv("MASTER_PORT", "29574")
monkeypatch.setenv("RANK", "0")
monkeypatch.setenv("WORLD_SIZE", "1")
monkeypatch.setenv("LOCAL_RANK", "0")
import fastvideo.models.dits.minimax_h3 as h3
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
from fastvideo.forward_context import set_forward_context
calls = dict.fromkeys(("rmsnorm_modulate", "residual_gate_rmsnorm_modulate", "qknorm_rope", "swiglu"), 0)
def _counting(name: str, real):
def wrapper(*args, **kwargs):
calls[name] += 1
return real(*args, **kwargs)
return wrapper
monkeypatch.setattr(h3, "fused_rmsnorm_modulate", _counting("rmsnorm_modulate", h3.fused_rmsnorm_modulate))
monkeypatch.setattr(h3, "fused_residual_gate_rmsnorm_modulate",
_counting("residual_gate_rmsnorm_modulate", h3.fused_residual_gate_rmsnorm_modulate))
monkeypatch.setattr(h3, "fused_qknorm_rope", _counting("qknorm_rope", h3.fused_qknorm_rope))
monkeypatch.setattr(h3, "minimax_h3_swiglu", _counting("swiglu", h3.minimax_h3_swiglu))
maybe_init_distributed_environment_and_model_parallel(1, 1)
try:
block = h3.MiniMaxH3TransformerBlock(
hidden_size=128,
num_attention_heads=1,
attention_head_dim=128,
ffn_dim=256,
time_embed_dim=64,
norm_eps=1e-5,
qk_norm_eps=1e-5,
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA, ),
quant_config=None,
prefix="minimax_h3.engagement_block",
fuse_modulate=True,
fuse_qknorm_rope=True,
fuse_swiglu=True,
)
device = torch.device("cuda")
block = block.to(device=device, dtype=torch.bfloat16).eval()
generator = torch.Generator(device=device).manual_seed(2026)
hidden_states = torch.randn(2, 12, 128, generator=generator, device=device, dtype=torch.bfloat16)
temb = torch.randn(2, 64, generator=generator, device=device, dtype=torch.bfloat16)
adaln_indices = torch.arange(12, device=device, dtype=torch.long).remainder(6)
position_ids = torch.zeros(12, 3, device=device, dtype=torch.float32)
position_ids[:, 0] = torch.arange(12, device=device)
rotary_emb = h3.MiniMaxH3RotaryPosEmbed(rope_freq_dim=16, rope_theta=10000.0).to(device)(position_ids)
inputs = dict(
hidden_states=hidden_states,
temb=temb,
adaln_indices=adaln_indices,
rotary_emb=tuple(value.to(torch.bfloat16) for value in rotary_emb),
original_seq_len=12,
)
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
block(**inputs)
engaged = dict(calls)
assert engaged == {
"rmsnorm_modulate": 1,
"residual_gate_rmsnorm_modulate": 1,
"qknorm_rope": 2,
"swiglu": 1,
}, engaged
grad_inputs = {**inputs, "hidden_states": hidden_states.clone().requires_grad_(True)}
with set_forward_context(current_timestep=0, attn_metadata=None):
block(**grad_inputs)
assert dict(calls) == engaged, f"a fusion ran under grad: {calls} vs {engaged}"
finally:
cleanup_dist_env_and_memory()
@@ -0,0 +1,173 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
import torch.nn.functional as F
from fastvideo.models.dits.minimax_h3_fusions.modulation import (
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
)
EPS = 1e-6
SOL_ENGINE_BF16_TOLERANCE = 3e-2
def _chunk_tables(rows: int, hidden_size: int, *, device: torch.device | str = "cpu", dtype=torch.float32):
wide = torch.randn(rows, 6 * hidden_size, device=device, dtype=dtype)
tables = wide.chunk(6, dim=-1)
assert all(table.stride() == (6 * hidden_size, 1) for table in tables)
assert all(not table.is_contiguous() for table in tables)
return tables
def _eager_rmsnorm_modulate(x, weight, scale, shift, index):
normed = F.rms_norm(x, (x.shape[-1], ), weight, EPS)
return normed * (1.0 + scale.index_select(0, index)) + shift.index_select(0, index)
def _eager_residual_gate_rmsnorm_modulate(residual, branch, gate, weight, scale, shift, index):
hidden = residual + gate.index_select(0, index) * branch
normed = F.rms_norm(hidden, (hidden.shape[-1], ), weight, EPS)
modulated = normed * (1.0 + scale.index_select(0, index)) + shift.index_select(0, index)
return hidden, modulated
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a CUDA GPU")
def test_bf16_sol_engine_fusions_match_production_eager_within_tolerance():
"""Sol-Engine keeps fused intermediates in FP32 until its output stores."""
pytest.importorskip("triton")
torch.manual_seed(2)
device = torch.device("cuda")
batch, sequence_length, hidden_size, table_rows = 2, 9, 5376, 6
residual = torch.randn(batch, sequence_length, hidden_size, device=device, dtype=torch.bfloat16)
branch = torch.randn_like(residual)
weight = torch.randn(hidden_size, device=device, dtype=torch.bfloat16)
shift, scale, gate, shift_mlp, scale_mlp, _ = _chunk_tables(
table_rows,
hidden_size,
device=device,
dtype=torch.bfloat16,
)
index = torch.tensor([5, 0, 4, 1, 3, 2, 5, 1, 0], device=device, dtype=torch.int64)
expected_norm1 = _eager_rmsnorm_modulate(residual, weight, scale, shift, index)
actual_norm1 = fused_rmsnorm_modulate(residual, weight, scale, shift, index, EPS)
expected_hidden, expected_norm2 = _eager_residual_gate_rmsnorm_modulate(
residual,
branch,
gate,
weight,
scale_mlp,
shift_mlp,
index,
)
actual_hidden, actual_norm2 = fused_residual_gate_rmsnorm_modulate(
residual,
branch,
gate,
weight,
scale_mlp,
shift_mlp,
index,
EPS,
)
# Production eager materializes BF16 after each PyTorch operator. The
# single-kernel Sol-Engine path deliberately removes those round points.
torch.testing.assert_close(
actual_norm1,
expected_norm1,
rtol=SOL_ENGINE_BF16_TOLERANCE,
atol=SOL_ENGINE_BF16_TOLERANCE,
)
torch.testing.assert_close(
actual_hidden,
expected_hidden,
rtol=SOL_ENGINE_BF16_TOLERANCE,
atol=SOL_ENGINE_BF16_TOLERANCE,
)
torch.testing.assert_close(
actual_norm2,
expected_norm2,
rtol=SOL_ENGINE_BF16_TOLERANCE,
atol=SOL_ENGINE_BF16_TOLERANCE,
)
@pytest.mark.parametrize(
("mutate", "error", "match"),
[
(lambda args: args | {"x": args["x"][0, 0]}, ValueError, "shape"),
(lambda args: args | {"weight": args["weight"][:-1]}, ValueError, "weight"),
(lambda args: args | {"index": args["index"][:-1]}, ValueError, "index"),
(lambda args: args | {"index": args["index"].float()}, TypeError, "index"),
(lambda args: args | {"eps": 0.0}, ValueError, "eps"),
],
)
def test_fusion_rejects_invalid_contracts(mutate, error, match):
args = {
"x": torch.randn(2, 3, 8),
"weight": torch.randn(8),
"scale": torch.randn(4, 8),
"shift": torch.randn(4, 8),
"index": torch.tensor([0, 3, 1]),
"eps": EPS,
}
with pytest.raises(error, match=match):
fused_rmsnorm_modulate(**mutate(args))
def test_residual_fusion_rejects_mismatched_branch():
residual = torch.randn(2, 3, 8)
with pytest.raises(ValueError, match="branch"):
fused_residual_gate_rmsnorm_modulate(
residual,
torch.randn(2, 2, 8),
torch.randn(4, 8),
torch.randn(8),
torch.randn(4, 8),
torch.randn(4, 8),
torch.tensor([0, 1, 2]),
EPS,
)
def test_triton_wrappers_fail_explicitly_on_cpu():
x = torch.randn(1, 2, 8)
branch = torch.randn_like(x)
weight = torch.randn(8)
gate = torch.randn(3, 8)
scale = torch.randn(3, 8)
shift = torch.randn(3, 8)
index = torch.tensor([0, 2])
with pytest.raises(RuntimeError, match="Triton|CUDA"):
fused_rmsnorm_modulate(x, weight, scale, shift, index, EPS)
with pytest.raises(RuntimeError, match="Triton|CUDA"):
fused_residual_gate_rmsnorm_modulate(x, branch, gate, weight, scale, shift, index, EPS)
def test_triton_wrappers_reject_autograd_before_backend_check():
x = torch.randn(1, 2, 8)
branch = torch.randn_like(x, requires_grad=True)
weight = torch.randn(8, requires_grad=True)
gate = torch.randn(3, 8)
scale = torch.randn(3, 8)
shift = torch.randn(3, 8)
index = torch.tensor([0, 2])
with pytest.raises(RuntimeError, match="forward-only"):
fused_rmsnorm_modulate(x, weight, scale, shift, index, EPS)
with pytest.raises(RuntimeError, match="forward-only"):
fused_residual_gate_rmsnorm_modulate(
x,
branch,
gate,
weight.detach(),
scale,
shift,
index,
EPS,
)
@@ -0,0 +1,192 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch
import torch.nn.functional as F
from fastvideo.models.dits.minimax_h3 import MiniMaxH3Attention
from fastvideo.models.dits.minimax_h3_fusions.qknorm_rope import (
HAVE_TRITON,
fused_qknorm_rope,
)
def _rotary_tables(
seq_len: int,
rotary_dim: int,
*,
dtype: torch.dtype,
device: torch.device | str,
) -> tuple[torch.Tensor, torch.Tensor]:
angles = torch.randn(seq_len, rotary_dim // 2, dtype=torch.float32, device=device)
angles = torch.cat((angles, angles), dim=-1)
return angles.cos().to(dtype), angles.sin().to(dtype)
def _eager_qknorm_rope(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
normalized = F.rms_norm(x, (x.shape[-1], ), weight, eps)
return MiniMaxH3Attention._apply_rotary_emb(normalized, (cos, sin))
def test_fused_qknorm_rope_rejects_invalid_rotary_dim() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
cos, sin = _rotary_tables(3, 130, dtype=x.dtype, device=x.device)
with pytest.raises(ValueError, match="rotary_dim must not exceed head_dim"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
cos = torch.randn(3, 95)
sin = torch.randn_like(cos)
with pytest.raises(ValueError, match="rotary_dim must be even"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
def test_fused_qknorm_rope_rejects_shape_mismatches() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
cos, sin = _rotary_tables(3, 96, dtype=x.dtype, device=x.device)
with pytest.raises(ValueError, match="x must have shape"):
fused_qknorm_rope(x[0], weight, cos, sin, 1e-6)
with pytest.raises(ValueError, match="weight must have shape"):
fused_qknorm_rope(x, weight[:-1], cos, sin, 1e-6)
with pytest.raises(ValueError, match="sequence length"):
fused_qknorm_rope(x, weight, cos[:-1], sin[:-1], 1e-6)
with pytest.raises(ValueError, match="sin must match cos shape"):
fused_qknorm_rope(x, weight, cos, sin[:, :-2], 1e-6)
@pytest.mark.parametrize("noncontiguous_input", ["weight", "cos", "sin"])
def test_fused_qknorm_rope_rejects_noncontiguous_linear_inputs(noncontiguous_input: str) -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(256)[::2]
cos = torch.randn(3, 192)[:, ::2]
sin = torch.randn(3, 192)[:, ::2]
assert not weight.is_contiguous()
assert not cos.is_contiguous()
assert not sin.is_contiguous()
if noncontiguous_input != "weight":
weight = weight.contiguous()
if noncontiguous_input != "cos":
cos = cos.contiguous()
if noncontiguous_input != "sin":
sin = sin.contiguous()
with pytest.raises(ValueError, match="must be contiguous"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
def test_fused_qknorm_rope_requires_matching_precast_dtype() -> None:
x = torch.randn(2, 3, 4, 128, dtype=torch.float32)
weight = torch.ones(128, dtype=torch.float32)
cos, sin = _rotary_tables(3, 96, dtype=torch.bfloat16, device=x.device)
with pytest.raises(TypeError, match="cos dtype must match x dtype"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
def test_fused_qknorm_rope_does_not_accept_missing_rotary_tables() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
with pytest.raises(TypeError, match="cos must be a torch.Tensor"):
fused_qknorm_rope(x, weight, None, None, 1e-6) # type: ignore[arg-type]
def test_fused_qknorm_rope_requires_cuda() -> None:
x = torch.randn(2, 3, 4, 128)
weight = torch.ones(128)
cos, sin = _rotary_tables(3, 96, dtype=x.dtype, device=x.device)
with pytest.raises(RuntimeError, match="requires CUDA"):
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
@pytest.mark.parametrize(
"rotary_dim,shape,use_input_view",
[
pytest.param(96, (2, 11, 5, 128), False, id="partial-96-batch2-seq11-heads5"),
pytest.param(128, (2, 7, 3, 128), True, id="full-128-noncontiguous-input-view"),
],
)
def test_fused_qknorm_rope_matches_eager_bf16_cuda(
rotary_dim: int,
shape: tuple[int, ...],
use_input_view: bool,
) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required for the Triton fusion")
if not HAVE_TRITON:
pytest.skip("Triton is required for the fusion")
torch.manual_seed(1)
device = torch.device("cuda")
if use_input_view:
x = torch.randn(*shape[:-1], shape[-1] * 2, dtype=torch.bfloat16, device=device)[..., ::2]
assert not x.is_contiguous()
else:
x = torch.randn(shape, dtype=torch.bfloat16, device=device)
weight = (1.0 + 0.05 * torch.randn(shape[-1], dtype=torch.bfloat16, device=device)).contiguous()
cos, sin = _rotary_tables(shape[1], rotary_dim, dtype=x.dtype, device=device)
with torch.inference_mode():
actual = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
expected = _eager_qknorm_rope(x, weight, cos, sin, 1e-6)
# The fused kernel keeps RMSNorm and both RoPE products in FP32 registers
# until its final BF16 store. Eager materializes BF16 intermediates, and
# PyTorch/Triton reductions need not use the same summation order.
torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2)
assert actual.shape == x.shape
assert actual.dtype == x.dtype
assert actual.is_contiguous()
def test_fused_qknorm_rope_matches_eager_beyond_int32_element_count() -> None:
"""Regression: kernel row offsets must be int64.
With int32 offsets, ``row * head_dim`` wraps once the flattened input
crosses 2**31 elements and the kernel reads/writes out of bounds (CUDA
illegal memory access). ``(1, 8_500_000, 2, 128)`` is 2.176e9 elements,
just past the boundary; for H3's 56 heads x 128 head_dim the equivalent
is ``batch*seq >= 299_593`` tokens per rank, reachable at SP=1.
GPU assumption: needs ~16 GiB free CUDA memory (input + output at
bf16 plus the fp32 rotary-table construction); skips below 20 GiB.
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required for the Triton fusion")
if not HAVE_TRITON:
pytest.skip("Triton is required for the fusion")
free_bytes, _ = torch.cuda.mem_get_info()
if free_bytes < 20 * 1024**3:
pytest.skip("needs ~20 GiB free GPU memory for a >2**31-element input")
heads, head_dim, rotary_dim = 2, 128, 96
seq_len = 8_500_000
assert seq_len * heads * head_dim > 2**31
torch.manual_seed(9)
device = torch.device("cuda")
x = torch.randn(1, seq_len, heads, head_dim, dtype=torch.bfloat16, device=device)
weight = (1.0 + 0.05 * torch.randn(head_dim, dtype=torch.bfloat16, device=device)).contiguous()
cos, sin = _rotary_tables(seq_len, rotary_dim, dtype=x.dtype, device=device)
with torch.inference_mode():
fused = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
torch.cuda.synchronize()
# Compare only head/tail slices against eager: a full-tensor eager
# reference would double peak memory for no extra coverage, and the
# tail rows are exactly the ones an int32 wrap corrupts first.
expected_head = _eager_qknorm_rope(x[:, :8], weight, cos[:8], sin[:8], 1e-6)
expected_tail = _eager_qknorm_rope(x[:, -8:], weight, cos[-8:], sin[-8:], 1e-6)
torch.testing.assert_close(fused[:, :8], expected_head, atol=2e-2, rtol=2e-2)
torch.testing.assert_close(fused[:, -8:], expected_tail, atol=2e-2, rtol=2e-2)
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Focused tests for MiniMax H3's value-first packed SwiGLU fusion."""
from __future__ import annotations
import pytest
import torch
import torch.nn.functional as F
from fastvideo.models.dits.minimax_h3_fusions.swiglu import minimax_h3_swiglu
def _require_bf16_triton_cuda() -> None:
if not torch.cuda.is_available():
pytest.skip("MiniMax H3 fused SwiGLU requires CUDA")
if not torch.cuda.is_bf16_supported():
pytest.skip("MiniMax H3 fused SwiGLU parity requires BF16 support")
pytest.importorskip("triton", reason="MiniMax H3 fused SwiGLU requires Triton")
def _assert_bf16_parity(x: torch.Tensor) -> None:
value, gate = x.chunk(2, dim=-1)
expected = value * F.silu(gate)
actual = minimax_h3_swiglu(x)
assert actual.shape == (*x.shape[:-1], x.shape[-1] // 2)
assert actual.dtype == x.dtype
assert actual.device == x.device
# Sol-Engine keeps the full SwiGLU expression in FP32 until the output
# store, while eager F.silu materializes a BF16 intermediate.
torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2)
@pytest.mark.parametrize("last_dim", [0, 7])
def test_minimax_h3_swiglu_rejects_nonpositive_or_odd_last_dimension(last_dim: int) -> None:
x = torch.empty((2, last_dim), dtype=torch.float32)
with pytest.raises(ValueError, match="positive even last dimension"):
minimax_h3_swiglu(x)
def test_minimax_h3_swiglu_strict_wrapper_rejects_cpu() -> None:
with pytest.raises(ValueError, match="requires a CUDA tensor"):
minimax_h3_swiglu(torch.randn(2, 8))
@pytest.mark.gpu
def test_minimax_h3_swiglu_multidimensional_bf16_gpu_parity() -> None:
_require_bf16_triton_cuda()
torch.manual_seed(0)
x = torch.randn((2, 3, 5, 66), device="cuda", dtype=torch.bfloat16)
_assert_bf16_parity(x)
@pytest.mark.gpu
def test_minimax_h3_swiglu_noncontiguous_bf16_gpu_parity() -> None:
_require_bf16_triton_cuda()
torch.manual_seed(1)
storage = torch.randn((2, 3, 148), device="cuda", dtype=torch.bfloat16)
x = storage[..., ::2]
assert not x.is_contiguous()
_assert_bf16_parity(x)
@pytest.mark.gpu
def test_minimax_h3_swiglu_real_ffn_dim_bf16_gpu() -> None:
_require_bf16_triton_cuda()
torch.manual_seed(2)
ffn_dim = 14336
x = torch.randn((1, 2 * ffn_dim), device="cuda", dtype=torch.bfloat16)
_assert_bf16_parity(x)
@@ -0,0 +1,428 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU contract tests for MiniMax H3 VAE compilation and profiling ranges.
The final test is a CUDA regression gate for the reduce-overhead tile path
(real tiled decode/encode with an unmocked ``_stitch_tiles``).
"""
from contextlib import contextmanager
from types import MethodType, SimpleNamespace
from typing import Any
from unittest.mock import Mock, patch
import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.models.vaes.minimax_h3_audio import (
MiniMaxH3AudioBigVGANDecoder,
MiniMaxH3AudioVAE,
)
from fastvideo.models.vaes.minimax_h3_video import (
AutoencoderKLMiniMaxH3,
MiniMaxH3VideoAttention,
MiniMaxH3VideoViTDecoder3d,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
def _empty_typed_module(module_type: type[nn.Module]) -> nn.Module:
"""Create a weightless instance that retains its production module type."""
module = object.__new__(module_type)
nn.Module.__init__(module)
return module
def _assert_dynamic_compile_selects_decoder(
vae_type: type[nn.Module],
decoder_type: type[nn.Module],
) -> None:
"""Verify one H3 VAE compiles only its top-level decoder in place."""
vae = _empty_typed_module(vae_type)
decoder = _empty_typed_module(decoder_type)
same_type_under_another_name = _empty_typed_module(decoder_type)
unrelated_submodule = nn.Identity()
vae.decoder = decoder
vae.same_type_under_another_name = same_type_under_another_name
vae.unrelated_submodule = unrelated_submodule
compiled_forward = Mock(name="compiled_forward")
compile_kwargs = {"backend": "inductor", "dynamic": False}
with patch(
"fastvideo.pipelines.composed_pipeline_base.torch.compile",
return_value=compiled_forward,
) as compile_mock:
compiled_count = ComposedPipelineBase._compile_with_conditions(vae, compile_kwargs)
assert compiled_count == 1
compile_mock.assert_called_once()
selected_forward = compile_mock.call_args.args[0]
assert selected_forward.__self__ is decoder
assert selected_forward.__func__ is decoder_type.forward
assert compile_mock.call_args.kwargs == compile_kwargs
assert decoder.forward is compiled_forward
assert "forward" not in same_type_under_another_name.__dict__
assert "forward" not in unrelated_submodule.__dict__
wrong_type_vae = _empty_typed_module(vae_type)
wrong_type_vae.decoder = nn.Identity()
with patch("fastvideo.pipelines.composed_pipeline_base.torch.compile") as wrong_type_compile:
wrong_type_count = ComposedPipelineBase._compile_with_conditions(wrong_type_vae, compile_kwargs)
assert wrong_type_count == 0
wrong_type_compile.assert_not_called()
def _assert_reduce_overhead_compile(compiled_function: Any) -> None:
"""Verify a class-owned compile boundary enables CUDA Graph replay."""
assert hasattr(compiled_function, "get_compiler_config")
assert compiled_function.get_compiler_config()["triton.cudagraphs"] is True
def test_video_attention_uses_selected_fastvideo_backend() -> None:
"""Pass BSHD tensors to the selected dense backend without forward metadata."""
backend_call: dict[str, Any] = {}
class RecordingAttentionImpl:
"""Record the backend construction and forward contracts."""
def __init__(self, **kwargs: Any) -> None:
backend_call["init"] = kwargs
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attention_metadata: Any,
) -> torch.Tensor:
"""Return values unchanged after recording the backend inputs."""
backend_call["shapes"] = (query.shape, key.shape, value.shape)
backend_call["metadata"] = attention_metadata
return value
class RecordingAttentionBackend:
"""Supply the recording implementation through the backend API."""
@staticmethod
def get_impl_cls() -> type[RecordingAttentionImpl]:
return RecordingAttentionImpl
with (
patch("fastvideo.platforms.current_platform") as current_platform,
patch(
"fastvideo.models.vaes.minimax_h3_video.get_attn_backend",
return_value=RecordingAttentionBackend,
) as get_backend,
):
current_platform.is_cuda_alike.return_value = True
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
attention.to_q = nn.Identity()
attention.to_k = nn.Identity()
attention.to_v = nn.Identity()
attention.norm_q = nn.Identity()
attention.norm_k = nn.Identity()
attention.to_out[0] = nn.Identity()
output = attention(torch.empty((1, 3, 8), device="meta"))
get_backend.assert_called_once_with(
4,
torch.bfloat16,
supported_attention_backends=(
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
),
)
assert backend_call["init"] == {
"num_heads": 2,
"head_size": 4,
"softmax_scale": 0.5,
"num_kv_heads": 2,
"causal": False,
}
assert backend_call["shapes"] == ((1, 3, 2, 4), ) * 3
assert backend_call["metadata"] is None
assert output.shape == (1, 3, 8)
def test_video_attention_cpu_uses_torch_sdpa() -> None:
"""Use PyTorch SDPA when H3 VAE attention receives CPU tensors."""
with (
patch("fastvideo.platforms.current_platform") as current_platform,
patch("fastvideo.models.vaes.minimax_h3_video.get_attn_backend") as get_backend,
):
current_platform.is_cuda_alike.return_value = False
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
get_backend.assert_not_called()
assert attention.attn_impl is None
attention.to_q = nn.Identity()
attention.to_k = nn.Identity()
attention.to_v = nn.Identity()
attention.norm_q = nn.Identity()
attention.norm_k = nn.Identity()
attention.to_out[0] = nn.Identity()
hidden_states = torch.randn(1, 3, 8)
query = hidden_states.unflatten(2, (2, 4)).permute(0, 2, 1, 3)
expected = F.scaled_dot_product_attention(query, query, query).permute(0, 2, 1, 3).flatten(2, 3)
torch.testing.assert_close(attention(hidden_states), expected)
def test_compile_with_conditions_selects_minimax_h3_video_decoder() -> None:
"""Compile the registered video decoder with the VAE runtime kwargs."""
assert not hasattr(MiniMaxH3VideoViTDecoder3d.forward, "get_compiler_config")
_assert_dynamic_compile_selects_decoder(AutoencoderKLMiniMaxH3, MiniMaxH3VideoViTDecoder3d)
def test_project_decoder_tile_uses_reduce_overhead_compile() -> None:
"""Compile the per-tile decoder-input projection with CUDA Graph replay."""
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._project_decoder_tile)
def test_stitch_tiles_uses_reduce_overhead_compile() -> None:
"""Compile spatial tile blending and concatenation with CUDA Graph replay."""
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._stitch_tiles)
def test_compile_with_conditions_selects_minimax_h3_audio_decoder() -> None:
"""Compile the audio VAE decoder that the H3 waveform decode path calls."""
_assert_dynamic_compile_selects_decoder(MiniMaxH3AudioVAE, MiniMaxH3AudioBigVGANDecoder)
def test_decode_emits_indexed_temporal_chunk_ranges() -> None:
"""Nest frame-segment ranges under each temporal decoder chunk range."""
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
vae.tokens_chunk_size = 1
vae.token_overlap = 1
vae.temporal_compression_ratio = 1
vae.frame_pre_padding = 0
vae.frame_overlap = 1
vae.config = SimpleNamespace(token_drop=1)
vae._decode_clip = Mock(return_value=torch.zeros((1, 1, 2, 1, 1)))
range_events = []
@contextmanager
def record_range(name: str):
range_events.append(("enter", name))
try:
yield
finally:
range_events.append(("exit", name))
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
decoded = vae._decode(torch.zeros((1, 1, 2, 1, 1)))
assert decoded.shape == (1, 1, 3, 1, 1)
assert vae._decode_clip.call_count == 2
assert range_events == [
("enter", "minimax_h3.vae.temporal_chunk.0"),
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.0"),
("enter", "minimax_h3.vae.temporal_chunk.1"),
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
("exit", "minimax_h3.vae.temporal_chunk.1"),
]
def test_decode_clip_no_spatial_tiling_stage_ranges() -> None:
"""Separate untiled latent projection and decoder ranges."""
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
vae.use_tiling = False
range_events = []
vae.post_quant_conv = nn.Identity()
vae.post_quant_conv.register_forward_hook(
lambda _module, _args, _output: range_events.append(("call", "post_quant_conv")))
vae.decoder = nn.Identity()
vae.decoder.register_forward_hook(
lambda _module, _args, _output: range_events.append(("call", "decoder_forward")))
latent_clip = torch.zeros((1, 1, 1, 2, 2))
@contextmanager
def record_range(name: str):
range_events.append(("enter", name))
try:
yield
finally:
range_events.append(("exit", name))
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
decoded_clip = vae._decode_clip(latent_clip)
assert decoded_clip is latent_clip
assert range_events == [
("enter", "minimax_h3.vae.decode_clip"),
("enter", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
("call", "post_quant_conv"),
("exit", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
("enter", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
("call", "decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip"),
]
def test_decode_clip_emits_tiled_stage_ranges() -> None:
"""Nest indexed decoder tiles between tile-splitting and stitching ranges."""
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
vae.use_tiling = True
vae.spatial_compression_ratio = 1
vae.tile_sample_min_height = 1
vae.tile_sample_min_width = 1
vae.tile_sample_min_overlap_height = 0
vae.tile_sample_min_overlap_width = 0
vae._split_tiles = Mock(side_effect=[
([0, 1], [1, 1], [0]),
([0, 1], [1, 1], [0]),
])
vae.post_quant_conv = nn.Identity()
vae._project_decoder_tile = Mock(side_effect=vae.post_quant_conv)
vae.decoder = nn.Identity()
stitched_clip = torch.zeros((1, 1, 1, 2, 2))
vae._stitch_tiles = Mock(return_value=stitched_clip)
range_events = []
@contextmanager
def record_range(name: str):
range_events.append(("enter", name))
try:
yield
finally:
range_events.append(("exit", name))
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
decoded_clip = vae._decode_clip(torch.zeros((1, 1, 1, 2, 2)))
# The tile driver must hand back a caller-owned copy: under
# mode="reduce-overhead" the stitched canvas is CUDA-graph pooled storage
# that the next replay overwrites, so returning it by identity is a bug.
assert decoded_clip is not stitched_clip
assert torch.equal(decoded_clip, stitched_clip)
assert vae._split_tiles.call_count == 2
assert vae._project_decoder_tile.call_count == 4
assert vae._stitch_tiles.call_count == 1
assert range_events == [
("enter", "minimax_h3.vae.decode_clip"),
("enter", "minimax_h3.vae.decode_clip.split_tiles"),
("exit", "minimax_h3.vae.decode_clip.split_tiles"),
("enter", "minimax_h3.vae.decode_clip.decode_tiles"),
("enter", "minimax_h3.vae.decode_clip.tile.0.0"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.0.0"),
("enter", "minimax_h3.vae.decode_clip.tile.0.1"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.0.1"),
("enter", "minimax_h3.vae.decode_clip.tile.1.0"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.1.0"),
("enter", "minimax_h3.vae.decode_clip.tile.1.1"),
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
("exit", "minimax_h3.vae.decode_clip.tile.1.1"),
("exit", "minimax_h3.vae.decode_clip.decode_tiles"),
("enter", "minimax_h3.vae.decode_clip.stitch_tiles"),
("exit", "minimax_h3.vae.decode_clip.stitch_tiles"),
("exit", "minimax_h3.vae.decode_clip"),
]
def _tiny_real_vae() -> AutoencoderKLMiniMaxH3:
"""Random-weight VAE small enough for a real tiled decode/encode on GPU."""
from fastvideo.configs.models.vaes.minimax_h3_video import (
MiniMaxH3VideoVAEArchConfig,
MiniMaxH3VideoVAEConfig,
)
arch = MiniMaxH3VideoVAEArchConfig(
latent_channels=4,
block_out_channels=(32, 32),
layers_per_block=1,
spatial_downsample_factors=(2, 2),
temporal_downsample_factors=(2, 2),
decoder_num_layers=1,
decoder_num_attention_heads=1,
decoder_attention_head_dim=8,
decoder_num_register_tokens=2,
decoder_ffn_mult=1,
latents_mean=(0.0, ) * 4,
latents_std=(1.0, ) * 4,
)
return AutoencoderKLMiniMaxH3(
MiniMaxH3VideoVAEConfig(
arch_config=arch,
use_tiling=False,
use_temporal_tiling=False,
use_parallel_tiling=False,
)).eval()
def _dynamo_original(compiled_function: Any) -> Any:
"""Return the eager callable behind a ``torch.compile``-decorated function."""
original = getattr(compiled_function, "_torchdynamo_orig_callable", None)
if original is None:
original = getattr(compiled_function, "__wrapped__", None)
assert original is not None, "cannot recover the eager tile helpers"
return original
@pytest.mark.skipif(not torch.cuda.is_available(), reason="reduce-overhead tile compile requires CUDA graphs")
@torch.inference_mode()
def test_tiled_decode_and_encode_survive_cudagraph_buffer_reuse_on_cuda() -> None:
"""Real tiled decode()/encode() with an unmocked reduce-overhead ``_stitch_tiles``.
Regression gate for the CUDA-graph output-clobbering bug: the stitched
canvas is a cudagraph static buffer, and the collect-then-``torch.cat``
consumers (``_decode``/``_encode``/``_encode_pixels``) hold chunk/clip
results across subsequent ``_stitch_tiles`` replays. Without the eager
``.clone()`` at the tile-driver returns, the first tiled ``decode()`` with
>=2 temporal chunks raises ``accessing tensor output of CUDAGraphs that
has been overwritten by a subsequent run``. This test needs >=2 chunks
(decode), >=2 clips (encode), and a >=2x2 spatial tile grid.
"""
torch.manual_seed(20260821)
vae = _tiny_real_vae().to("cuda")
vae.enable_tiling(16, 16, 4, 4)
# 8 latent tokens = 2 temporal chunks (tokens_chunk_size 5); 8x8 latents =
# 32x32 pixels = a 2x2 grid of 16px tiles.
z = torch.randn(1, 4, 8, 8, 8, device="cuda")
pad_tokens, num_chunks, _ = vae._temporal_decode_plan(z.shape[2])
assert num_chunks >= 2, "decode workload must span multiple stitch replays"
decoded_first = vae.decode(z).sample
decoded_second = vae.decode(z).sample
assert torch.equal(decoded_first, decoded_second)
# 34 frames = 2 encode clips of clip_length 17 -> 2 stitch replays.
pixels = torch.rand(1, 3, 34, 32, 32, device="cuda")
encoded_first = vae.encode(pixels).latent_dist.parameters
encoded_second = vae.encode(pixels).latent_dist.parameters
assert torch.equal(encoded_first, encoded_second)
uint8_pixels = torch.randint(0, 256, (1, 3, 34, 32, 32), dtype=torch.uint8)
streamed_first = vae.encode_pixels(uint8_pixels).latent_dist.parameters
streamed_second = vae.encode_pixels(uint8_pixels).latent_dist.parameters
assert torch.equal(streamed_first, streamed_second)
# Output parity vs the fully eager tile helpers (same weights, same math;
# the tolerance absorbs inductor fusion reassociation only).
eager_vae = _tiny_real_vae().to("cuda")
eager_vae.load_state_dict(vae.state_dict())
eager_vae.enable_tiling(16, 16, 4, 4)
eager_vae._stitch_tiles = MethodType(_dynamo_original(AutoencoderKLMiniMaxH3._stitch_tiles), eager_vae)
eager_vae._project_decoder_tile = MethodType(_dynamo_original(AutoencoderKLMiniMaxH3._project_decoder_tile),
eager_vae)
torch.testing.assert_close(decoded_first, eager_vae.decode(z).sample, atol=2e-4, rtol=2e-4)
torch.testing.assert_close(encoded_first, eager_vae.encode(pixels).latent_dist.parameters, atol=2e-4, rtol=2e-4)
+30 -1
View File
@@ -1,11 +1,40 @@
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines import ForwardBatch
from fastvideo.worker.gpu_worker import Worker
from fastvideo.worker.gpu_worker import Worker, _log_cuda_device_uuid
def test_cuda_device_uuid_receipt_is_disabled_without_nvtx_profiling(monkeypatch) -> None:
"""Avoid NVIDIA property access during ordinary worker initialization."""
get_device_properties = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
monkeypatch.setattr(torch.cuda, "get_device_properties", get_device_properties)
_log_cuda_device_uuid(0, torch.device("cuda:0"))
get_device_properties.assert_not_called()
def test_cuda_device_uuid_receipt_identifies_profiled_worker(monkeypatch) -> None:
"""Bind one profiled worker rank to its NVIDIA device UUID in logs."""
log_info = Mock()
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
monkeypatch.setattr(torch.cuda, "get_device_properties", lambda device: SimpleNamespace(uuid="device-uuid"))
monkeypatch.setattr("fastvideo.worker.gpu_worker.logger.info", log_info)
_log_cuda_device_uuid(2, torch.device("cuda:0"))
log_info.assert_called_once_with(
"Worker %d CUDA device UUID: GPU-%s",
2,
"device-uuid",
local_main_process_only=False,
)
def _worker_returning(output_batch: ForwardBatch) -> Worker:
+11
View File
@@ -4,6 +4,7 @@ from typing import Any, cast
import torch
import fastvideo.envs as envs
from fastvideo.distributed import (cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel)
from fastvideo.distributed.parallel_state import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -13,6 +14,14 @@ from fastvideo.pipelines import ForwardBatch, LoRAPipeline, build_pipeline
logger = init_logger(__name__)
def _log_cuda_device_uuid(rank: int, device: torch.device) -> None:
"""Record an NVIDIA worker UUID when external NVTX profiling is enabled."""
if not envs.FASTVIDEO_NVTX_PROFILE:
return
device_uuid = torch.cuda.get_device_properties(device).uuid
logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False)
class Worker:
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str):
@@ -61,6 +70,8 @@ class Worker:
if current_platform.is_cuda_alike():
torch.cuda.set_device(self.device)
self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0]
if current_platform.is_cuda():
_log_cuda_device_uuid(self.rank, self.device)
else:
# For MPS, we can't get memory info the same way
self.init_gpu_memory = 0
@@ -6,11 +6,8 @@ pipeline with FastVideo's production ``TextEncoderLoader`` path. It covers
the three numerical branches the H3 pipelines exercise: text-only tokens,
image features, and video features.
The production encoder is built only as far as the layer-50 conditioning tap
by default (``num_hidden_layers_override``), so it returns fewer hidden states
than the official full stack. Every state it does build is compared
bit-exactly against the official value at the same index, which pins the tap
and would catch a truncated stack that still applied the final norm.
The production encoder returns only the selected layer-50 hidden state, which
is compared bit-exactly with the same state from the official full stack.
"""
from __future__ import annotations
@@ -152,13 +149,13 @@ def _make_cases(root: Path) -> dict[str, dict[str, torch.Tensor]]:
return cases
def _run_cases(
def _run_reference_cases(
model: torch.nn.Module,
cases: dict[str, dict[str, torch.Tensor]],
device: torch.device,
) -> dict[str, tuple[torch.Tensor, ...]]:
) -> dict[str, torch.Tensor]:
dtype = next(model.parameters()).dtype
outputs: dict[str, tuple[torch.Tensor, ...]] = {}
outputs: dict[str, torch.Tensor] = {}
for name, case in cases.items():
inputs = {
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
@@ -172,7 +169,28 @@ def _run_cases(
)
assert result.hidden_states is not None
assert len(result.hidden_states) > MINIMAX_H3_TEXT_ENCODER_LAYER
outputs[name] = tuple(hidden_state.detach().cpu() for hidden_state in result.hidden_states)
outputs[name] = result.hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER][0].detach().cpu()
return outputs
def _run_production_cases(
model: torch.nn.Module,
cases: dict[str, dict[str, torch.Tensor]],
device: torch.device,
) -> dict[str, torch.Tensor]:
dtype = next(model.parameters()).dtype
outputs: dict[str, torch.Tensor] = {}
for name, case in cases.items():
inputs = {
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
for key, value in case.items()
if key not in {"attention_mask", "mm_token_type_ids"}
}
inputs["input_ids"] = inputs["input_ids"][0]
with torch.inference_mode():
result = model(**inputs)
assert result.ndim == 2
outputs[name] = result.detach().cpu()
return outputs
@@ -217,32 +235,25 @@ def test_minimax_h3_qwen3_vl_parity() -> None:
assert not load_errors, f"Official Qwen3-VL checkpoint did not load strictly: {load_errors}"
official = official_full.model.eval().to(device)
del official_full
expected = _run_cases(official, cases, device)
expected = _run_reference_cases(official, cases, device)
del official
_reclaim_vram()
production = TextEncoderLoader().load(str(root / "text_encoder"), _production_loader_args())
assert getattr(production, "_fastvideo_input_device", device) == device
actual = _run_cases(production, cases, device)
actual = _run_production_cases(production, cases, device)
assert actual.keys() == expected.keys()
# The production stack is built only as far as the conditioning tap by
# default (``num_hidden_layers_override``), so it yields one hidden state
# per built layer plus the embeddings, while the official model always
# yields the full tuple. Every state the production model produces must be
# bit-identical to the official value at the same index; the shared-prefix
# comparison would in particular catch a truncated stack that still
# applied the final norm, which is the failure mode that silently changes
# conditioning. With the override set to None the lengths are equal and
# this remains the original full comparison, final normed state included.
built_layers = int(production.language_model.num_layers)
for name in expected:
assert len(actual[name]) == built_layers + 1
assert len(actual[name]) <= len(expected[name])
for layer, (result, reference) in enumerate(zip(actual[name], expected[name], strict=False)):
assert_close(result, reference, atol=0.0, rtol=0.0, msg=lambda message: f"{name} layer {layer}: {message}")
result = actual[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
reference = expected[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
result = actual[name]
reference = expected[name]
assert_close(
result,
reference,
atol=0.0,
rtol=0.0,
msg=lambda message: f"{name} layer {MINIMAX_H3_TEXT_ENCODER_LAYER}: {message}",
)
drift = (result.float() - reference.float()).abs()
print(
f"{name}: max_abs={drift.max().item():.8f} mean_abs={drift.mean().item():.8f}",
+3 -4
View File
@@ -57,10 +57,9 @@ pytest \
```
With a gate enabled, missing CUDA, source, or weights is a failure. Recorded component evidence is exact for both DiT
partitions, the video VAE, and all Qwen3-VL hidden states; audio decode has maximum absolute drift `2.4e-7`. The
production Qwen3-VL stack is now built only to the layer-50 conditioning tap by default
(`num_hidden_layers_override`), so the encoder gate compares every hidden state the production model builds
bit-exactly against the official full stack at the same index.
partitions and the video VAE; audio decode has maximum absolute drift `2.4e-7`. The encoder gate compares the slim
forward's selected layer-50 hidden state bit-exactly against the same state from the official full stack across text,
image, and video inputs.
The video VAE test verifies the reference checkout at commit
`abc5e9bf71fd38f53cd471bc3acaa84bc5ecbfdc` and compares the production CPU `uint8` `encode_pixels()` path against