diff --git a/.gitignore b/.gitignore index 13c5fbc871..88035e8c2b 100644 --- a/.gitignore +++ b/.gitignore @@ -142,3 +142,4 @@ fastvideo/tests/ssim/reference_videos/** *.nvimlog .nvimlog .python-version +scripts/benchmarks/minimax_h3_pro6000/headline_results/ diff --git a/docs/assets/cookbook-recipes.json b/docs/assets/cookbook-recipes.json index e89e58d121..5f26fe1bb0 100644 --- a/docs/assets/cookbook-recipes.json +++ b/docs/assets/cookbook-recipes.json @@ -1,5 +1,5 @@ { - "version": 12, + "version": 14, "recipes": [ { "id": "fastwan21-t2v", @@ -601,6 +601,77 @@ "Height, width, frames, and steps in the YAML are examples. Edit them or pass CLI flags. See docs/getting_started/installation/spark_pair.md." ] }, + { + "id": "compacth3-rtx5090", + "group": "compacth3-rtx5090", + "group_label": "CompactH3 on RTX 5090", + "group_task": "4-step 42-block text to video + audio", + "family": "minimax_h3", + "stage": "inference", + "task": "Few-step text to video (with audio)", + "label": "CompactH3 NVFP4 on RTX 5090", + "summary": "Run the 42-block CompactH3 NVFP4 DiT with the NVFP4 Qwen3-VL encoder, Comfy int8-convrot VAE, and SageAttention3 FP4 on one 32 GB RTX 5090. Sequential load parks the encoder in pinned host RAM.", + "model": "./CompactH3", + "source": "examples/inference/basic/basic_compacth3_rtx5090.yaml", + "serving": { + "source": "examples/serving/openai_compacth3_rtx5090.yaml", + "install": "UV_TORCH_BACKEND=cu128 uv pip install -e \".[fasth3]\"", + "prepare": "hf download FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree --local-dir ./CompactH3 --include model_index.json --include \"tokenizer/**\" --include \"processor/**\" --include \"scheduler/**\" --include \"audio_scheduler/**\" --include \"audio_vae/**\" --include \"vae/**\"\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-bf16 --local-dir ./CompactH3/transformer\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-nvfp4 --local-dir ./CompactH3/transformer --include nvfp4_weights.safetensors\nhf download KyleNeverGivesUp/FastH3-text-encoder-nvfp4 --local-dir ./CompactH3/text_encoder\nhf download Comfy-Org/MiniMax-H3 --local-dir ./Comfy-MiniMax-H3 --include vae/minimax_h3_video_vae_int8_convrot.safetensors\ncp ./Comfy-MiniMax-H3/vae/minimax_h3_video_vae_int8_convrot.safetensors ./CompactH3/vae/", + "env": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a" + }, + "command": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 fastvideo generate --config examples/inference/basic/basic_compacth3_rtx5090.yaml", + "gpu_types": ["NVIDIA"], + "hardware": { + "platform": "cuda", + "gpu_count": 1, + "evidence": "source-configured" + }, + "evidence": "Source-backed", + "expected_artifact": "MP4 under outputs/compacth3_rtx5090/", + "modes": ["T2VA", "CompactH3 NVFP4", "RTX 5090"], + "limitations": [ + "Assemble ./CompactH3 before running. The DiT NVFP4 export and encoder snapshot are gated; run huggingface-cli login and accept each repo license.", + "32 GB cannot keep the NVFP4 encoder and DiT on the GPU together. Keep h3_sequential_load on and lazy_module_load off.", + "Blackwell sm_120 needs ATTN_QAT_INFER, FLASHINFER_CUDA_ARCH_LIST=12.0a, and a CUDA 12.8 PyTorch wheel.", + "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768; 1344x768 on 32 GB is unmeasured." + ] + }, + { + "id": "compacth3-rtx-pro6000", + "group": "compacth3-rtx-pro6000", + "group_label": "CompactH3 on RTX PRO 6000", + "group_task": "4-step 42-block text to video + audio", + "family": "minimax_h3", + "stage": "inference", + "task": "Few-step text to video (with audio)", + "label": "CompactH3 NVFP4 on RTX PRO 6000 Blackwell", + "summary": "Run CompactH3 NVFP4 with the encoder, DiT, and int8-convrot VAE resident on one 96 GB RTX PRO 6000 Blackwell. The checked-in example is 1344x768 and 124 frames (5.17 s) with VAE torch.compile. Use 832x480 for clip-queue playground traffic.", + "model": "./CompactH3", + "source": "examples/inference/basic/basic_compacth3_rtx_pro6000.yaml", + "serving": { + "source": "examples/serving/openai_compacth3_rtx_pro6000.yaml", + "install": "UV_TORCH_BACKEND=cu128 uv pip install -e \".[fasth3]\"", + "prepare": "hf download FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree --local-dir ./CompactH3 --include model_index.json --include \"tokenizer/**\" --include \"processor/**\" --include \"scheduler/**\" --include \"audio_scheduler/**\" --include \"audio_vae/**\" --include \"vae/**\"\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-bf16 --local-dir ./CompactH3/transformer\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-nvfp4 --local-dir ./CompactH3/transformer --include nvfp4_weights.safetensors\nhf download KyleNeverGivesUp/FastH3-text-encoder-nvfp4 --local-dir ./CompactH3/text_encoder\nhf download Comfy-Org/MiniMax-H3 --local-dir ./Comfy-MiniMax-H3 --include vae/minimax_h3_video_vae_int8_convrot.safetensors\ncp ./Comfy-MiniMax-H3/vae/minimax_h3_video_vae_int8_convrot.safetensors ./CompactH3/vae/", + "env": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a" + }, + "command": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 fastvideo generate --config examples/inference/basic/basic_compacth3_rtx_pro6000.yaml", + "gpu_types": ["NVIDIA"], + "hardware": { + "platform": "cuda", + "gpu_count": 1, + "evidence": "source-configured" + }, + "evidence": "Source-backed", + "expected_artifact": "MP4 under outputs/compacth3_rtx_pro6000/", + "modes": ["T2VA", "CompactH3 NVFP4", "RTX PRO 6000"], + "limitations": [ + "Assemble ./CompactH3 before running. The DiT NVFP4 export and encoder snapshot are gated; run huggingface-cli login and accept each repo license.", + "96 GB keeps the encoder, DiT, and VAE on GPU. Do not enable h3_sequential_load or lazy_module_load on this box.", + "Enable compile.vae_enabled. Leave inference_torch_compile off: FlashInfer and Sage3 custom ops cannot be compiled.", + "Blackwell sm_120 needs ATTN_QAT_INFER, FLASHINFER_CUDA_ARCH_LIST=12.0a, and a CUDA 12.8 PyTorch wheel.", + "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768. Dense CompactH3 has zero VSA gates; the pipeline raises if VIDEO_SPARSE_ATTN_H3 is loaded with all-zero to_gate_compress weights." + ] + }, { "id": "fasth3-8step-v2-cuda", "group": "fasth3-8step-v2", diff --git a/docs/contributing/env_vars.md b/docs/contributing/env_vars.md index 4f47ddcb1f..883e2d8ac8 100644 --- a/docs/contributing/env_vars.md +++ b/docs/contributing/env_vars.md @@ -240,6 +240,35 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_VBENCH_FULL_INFO_JSON` | str | unset | eval | Path to VBench_full_info.json, used instead of the vendored copy. Deprecated names: `VBENCH_FULL_INFO_JSON`. | | `FASTVIDEO_FVD_REF_FEATURES` | str | unset | eval | Cached reference-feature file for the FVD metric. | | `FASTVIDEO_FAD_REF_FEATURES` | str | unset | eval | Cached reference-feature file for the audio Frechet distance metric. | +| `FASTVIDEO_H3_VSA_FP4` | bool | `0` | attention | Run MiniMax-H3 VSA attention on the block-sparse SageAttention3 FP4 kernel (sm_120, no-grad, single sequence-parallel rank). | +| `FASTVIDEO_H3_VSA_TILE_FIRST` | bool | `0` | attention | Single-rank MiniMax-H3 VSA with one tile gather of the block input instead of separate Q/K/V/gate scatters. | +| `FASTVIDEO_H3_VSA_SM89_KERNEL` | one of original, bf16, int8 | `original` | attention | Fine-attention kernel for MiniMax-H3 VSA on sm_89: original, bf16, or int8 (INT8 QK, BF16 PV). | +| `FASTVIDEO_H3_SIM_SP_FP8` | bool | `0` | debug | Simulate the FP8 sequence-parallel exchange of the MiniMax-H3 FP4 VSA path on one rank. | +| `FASTVIDEO_H3_FFN_CHUNK_TOKENS` | int | `0` | performance | Inference-only MiniMax-H3 FFN token chunk size; 0 runs the FFN unchunked. | +| `FASTVIDEO_H3_FP8_ATTENTION` | bool | `0` | performance | With NVFP4 layer_profile h3_dit_ffn, run MiniMax-H3 attention projections in FP8. | +| `FASTVIDEO_H3_FP8_GRANULARITY` | one of tensor, channel | `tensor` | performance | FP8 scaling granularity for FASTVIDEO_H3_FP8_ATTENTION. | +| `FASTVIDEO_NVFP4_MM_BACKEND` | str | `auto` | performance | FlashInfer mm_fp4 backend for NVFP4 linears, e.g. auto or cutlass. | +| `FASTVIDEO_NVFP4_ACT_AMAX` | path | unset | performance | JSON of calibrated NVFP4 input amax per linear, keyed b<block>.<sub> or full prefix; sets a static activation scale. | +| `FASTVIDEO_NVFP4_DYNAMIC_ACT` | str | `""` | performance | NVFP4 linears that derive the activation scale per call: all, or comma-separated layer-name suffixes such as ff.fc_out. | +| `FASTVIDEO_H3_ADALN_CACHE` | bool | `0` | performance | Cache MiniMax-H3 AdaLN modulation per timestep instead of keeping the projection weights resident. | +| `FASTVIDEO_H3_ADALN_TABLE` | path | unset | performance | Precomputed MiniMax-H3 AdaLN modulation table; enables the cache and skips loading the AdaLN projection weights. | +| `FASTVIDEO_H3_ADALN_DUMP` | path | unset | debug | Write the MiniMax-H3 AdaLN modulation table to this path while sampling. | +| `FASTVIDEO_H3_SPLICE_TRANSFORMER` | path | unset | eval | Second MiniMax-H3 transformer that runs the late DMD steps (checkpoint step-splice evaluation). | +| `FASTVIDEO_H3_SPLICE_FROM_STEP` | int | `4` | eval | First denoising step run by FASTVIDEO_H3_SPLICE_TRANSFORMER. | +| `FASTVIDEO_H3_ENCODER_LAYERWISE` | bool | `0` | performance | Stream MiniMax-H3 text-encoder language layers through exact-size pinned host memory (text-only prompts). | +| `FASTVIDEO_H3_ENCODER_FUSED_DEQUANT` | bool | `0` | performance | Expand the serialized NVFP4 MiniMax-H3 text encoder with one fused Triton pass on GPUs without FP4 GEMM. | +| `FASTVIDEO_H3_VAE_TILE_BATCH` | int | `1` | performance | Spatial tiles per MiniMax-H3 video VAE decoder call; 1 decodes per tile. | +| `FASTVIDEO_H3_VAE_INT8_SHARED_QKV` | bool | `0` | performance | Share the INT8 activation rotation and quantization across the MiniMax-H3 VAE Q/K/V projections. | +| `FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW` | bool | `0` | performance | Use transposed weight views in the MiniMax-H3 VAE INT8 projections. | +| `FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT` | bool | `0` | performance | Fused dequantization epilogue for the MiniMax-H3 VAE INT8 projections. | +| `FASTVIDEO_H3_PINNED_SWAP` | bool | `1` | performance | Swap offloaded MiniMax-H3 modules through exact-size pinned host arenas. | +| `FASTVIDEO_H3_PARK_MODULES` | str | unset | performance | Comma-separated MiniMax-H3 denoise modules parked on the host while the text encoder runs, e.g. vae,audio_vae. | +| `FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS` | bool | `0` | performance | Layerwise offload also streams large buffers such as packed FP4/FP8 weights. | +| `FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS` | int | `0` | performance | Keep the first N layerwise-offloaded blocks resident on the GPU. | +| `FASTVIDEO_H3_SP_PROFILE` | bool | `0` | profiling | CUDA-event spans per stage over one MiniMax-H3 FP4 VSA DiT forward. | +| `FASTVIDEO_H3_CAPTURE_QKV` | path | unset | debug | Directory for captured real MiniMax-H3 Q/K/V attention inputs. | +| `FASTVIDEO_CUDA_MEMORY_CAP_GIB` | float | `0.0` | debug | Cap this process's CUDA allocator at this many GiB to emulate a smaller GPU; 0 leaves it uncapped. | +| `FASTVIDEO_MEMORY_REPORT` | bool | `0` | debug | Log bytes held per pipeline component by device and dtype after loading. | | `FASTVIDEO_TEST_LTX2_OVERFIT_DATA_DIR` | str | `data/cats` | test | Raw data directory for preprocess_ltx2_overfit.py. Deprecated names: `LTX2_OVERFIT_DATA_DIR`. | | `FASTVIDEO_TEST_LTX2_OVERFIT_CAPTION_JSON` | str | `videos2caption_1_sample.json` | test | Caption file, relative to the raw data directory. Deprecated names: `LTX2_OVERFIT_CAPTION_JSON`. | | `FASTVIDEO_TEST_LTX2_OVERFIT_VIDEO_SUBDIR` | str | `video` | test | Video subdirectory, relative to the raw data directory. Deprecated names: `LTX2_OVERFIT_VIDEO_SUBDIR`. | diff --git a/docs/cookbook/cosmos.md b/docs/cookbook/cosmos.md index fbb813fae6..4d9c8c0fbc 100644 --- a/docs/cookbook/cosmos.md +++ b/docs/cookbook/cosmos.md @@ -5,7 +5,7 @@ hide: # Cosmos recipes -
+
All model families
diff --git a/docs/cookbook/flux.md b/docs/cookbook/flux.md index f2f159aed7..160c56ca58 100644 --- a/docs/cookbook/flux.md +++ b/docs/cookbook/flux.md @@ -5,7 +5,7 @@ hide: # FLUX recipes -
+
All model families
diff --git a/docs/cookbook/glm-image.md b/docs/cookbook/glm-image.md index ed88351a17..93f8b5cf2e 100644 --- a/docs/cookbook/glm-image.md +++ b/docs/cookbook/glm-image.md @@ -5,7 +5,7 @@ hide: # GLM-Image recipes -
+
All model families
diff --git a/docs/cookbook/hunyuan.md b/docs/cookbook/hunyuan.md index 66acf7f288..dede8af94f 100644 --- a/docs/cookbook/hunyuan.md +++ b/docs/cookbook/hunyuan.md @@ -5,7 +5,7 @@ hide: # Hunyuan recipes -
+
All model families
diff --git a/docs/cookbook/kandinsky5.md b/docs/cookbook/kandinsky5.md index 119296eae7..d899dd977a 100644 --- a/docs/cookbook/kandinsky5.md +++ b/docs/cookbook/kandinsky5.md @@ -5,7 +5,7 @@ hide: # Kandinsky 5 recipes -
+
All model families
diff --git a/docs/cookbook/longcat.md b/docs/cookbook/longcat.md index ce2a81574a..28f303deb3 100644 --- a/docs/cookbook/longcat.md +++ b/docs/cookbook/longcat.md @@ -5,7 +5,7 @@ hide: # LongCat recipes -
+
All model families
diff --git a/docs/cookbook/ltx.md b/docs/cookbook/ltx.md index a9354f30da..f751661c9d 100644 --- a/docs/cookbook/ltx.md +++ b/docs/cookbook/ltx.md @@ -5,7 +5,7 @@ hide: # LTX recipes -
+
All model families
diff --git a/docs/cookbook/matrix-game.md b/docs/cookbook/matrix-game.md index 26160e8650..b6d4e16cd8 100644 --- a/docs/cookbook/matrix-game.md +++ b/docs/cookbook/matrix-game.md @@ -5,7 +5,7 @@ hide: # Matrix Game recipes -
+
All model families
diff --git a/docs/cookbook/minimax-h3.md b/docs/cookbook/minimax-h3.md index 01f77a7121..e01ee210a4 100644 --- a/docs/cookbook/minimax-h3.md +++ b/docs/cookbook/minimax-h3.md @@ -10,8 +10,10 @@ launch. Some Hub repo names still say Preview. That name is historical. V1 is a full model, not a demo. **V2** is the eight-step checkpoint. More forwards is why V2 is the higher-quality FastH3. The V2 schedule contract is in [FastH3 distilled checkpoint schedules](../inference/fasth3-distilled.md). +**CompactH3** is the 42-block 20B NVFP4 H3 checkpoint for one Blackwell GPU +(RTX 5090 or RTX PRO 6000). The FastH3 V1 and V2 recipes are unchanged. -
+
All model families
@@ -21,8 +23,8 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in

Primary focus · Inference

MiniMax H3 recipes

-

Generate video and audio with H3. Run a server on CUDA, one DGX Spark, or Apple Silicon MLX to iterate on prompts, or call the pipeline directly from Python.

- 9 maintained recipes +

Generate video and audio with H3. Run a server on CUDA, one Blackwell GPU, one DGX Spark, or Apple Silicon MLX to iterate on prompts, or call the pipeline directly from Python.

+ 11 maintained recipes
@@ -47,7 +49,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in

Supported modes

CUDA covers T2VA, FL2VA, and Ref2VA on the full checkpoint, plus FastH3 - V1 and FastH3 V2. FastH3 V1 also has a DGX Spark runtime with + V1 and FastH3 V2, plus CompactH3 NVFP4 on one Blackwell GPU. FastH3 V1 also has a DGX Spark runtime with a 1-Spark or 2-Spark device row. MLX is T2VA only: V1 and V2. Temporal --fast, spatial --fast-spatial, and opt-in VSA are flags on the same MLX script, not extra recipes. @@ -64,7 +66,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in T2VA - Full H3, FastH3 V1, FastH3 LoRA, FastH3 V2 + Full H3, FastH3 V1, FastH3 LoRA, FastH3 V2, CompactH3 NVFP4 FastH3 V1 or FastH3 V2 after a local DiT conversion @@ -102,6 +104,11 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in FastH3 V1 on one GB10, or two Sparks with Ray sequence parallel (sp_size=2) over QSFP RoCE. Select NVIDIA DGX Spark, then 1 Spark or 2 Sparks. Not wired + + CompactH3 NVFP4 + 42-block 20B checkpoint on one RTX 5090 (32 GB, sequential encoder offload) or one RTX PRO 6000 Blackwell (96 GB, encoder+DiT+VAE resident). SageAttention3 FP4, packed NVFP4 DiT, Comfy int8-convrot VAE. + Not wired +

@@ -110,7 +117,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in

Pick an H3 recipe and runtime

-

Choose the result you want, then use a maintained CUDA, DGX Spark, or MLX path. +

Choose the result you want, then use a maintained CUDA, Blackwell, DGX Spark, or MLX path. Device claims stay tied to checked-in sources and recorded runs.

@@ -275,7 +282,8 @@ cd FastVideo
  • The MLX source runtime supports T2VA, optional temporal --fast, optional spatial --fast-spatial, and opt-in VSA on --include-vsa checkpoints. FastH3 V2 MLX converts with --include-vsa and runs eight forwards. FL2VA, Ref2VA, and two-pass refinement are not wired.
  • GPU count and VAE decode backend are configurable in the builder above for FastH3 CUDA recipes. Only the value shown by default has a recorded run; other supported values are unmeasured here.
  • DGX Spark is a runtime on FastH3 V1, not a separate family card. Select NVIDIA DGX Spark, then 1 Spark or 2 Sparks. The CUDA GPU-count knob does not apply to Spark.
  • -
  • GB10 has no FA4 / sm_100a VSA kernel. Keep FASTVIDEO_FA4=0 and FASTVIDEO_VSA_SM100A=0. Legal num_frames values are 17n+5, capped at 362 (15.08 s). A 345-frame request on one Spark can OOM.
  • +
  • CompactH3 NVFP4 is one Blackwell GPU. RTX 5090 (32 GB) parks the encoder in pinned host RAM. RTX PRO 6000 Blackwell (96 GB) keeps encoder, DiT, and VAE resident. Keep FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER, FASTVIDEO_FA4=0, FASTVIDEO_VSA_SM100A=0, and FLASHINFER_CUDA_ARCH_LIST=12.0a. On PRO 6000 enable VAE compile and leave DiT inference_torch_compile off. CompactH3 is a dense prune; do not enable VIDEO_SPARSE_ATTN_H3 until a VSA-trained student exists.
  • +
  • GB10 has no FA4 / sm_100a VSA kernel. Keep FASTVIDEO_FA4=0 and FASTVIDEO_VSA_SM100A=0. Legal num_frames values are 17n+5, capped at 362 (15.08 s). A 345-frame request on one Spark can OOM. Native 16:9 sizes include 832×480 and 1344×768.
  • Gated or missing checkpoints: run huggingface-cli login and confirm you accepted the model's license on Hugging Face.
  • diff --git a/docs/cookbook/mmaudio.md b/docs/cookbook/mmaudio.md index 5e127c9b2b..544241d44a 100644 --- a/docs/cookbook/mmaudio.md +++ b/docs/cookbook/mmaudio.md @@ -5,7 +5,7 @@ hide: # MMAudio recipes -
    +
    All model families
    diff --git a/docs/cookbook/stable-audio.md b/docs/cookbook/stable-audio.md index 20dc7bfd7a..5fa15b299f 100644 --- a/docs/cookbook/stable-audio.md +++ b/docs/cookbook/stable-audio.md @@ -5,7 +5,7 @@ hide: # Stable Audio recipes -
    +
    All model families
    diff --git a/docs/cookbook/stable-diffusion.md b/docs/cookbook/stable-diffusion.md index 07c87fdb59..845a305941 100644 --- a/docs/cookbook/stable-diffusion.md +++ b/docs/cookbook/stable-diffusion.md @@ -5,7 +5,7 @@ hide: # Stable Diffusion recipes -
    +
    All model families
    diff --git a/docs/cookbook/turbodiffusion.md b/docs/cookbook/turbodiffusion.md index 14820b5635..df3dac64f5 100644 --- a/docs/cookbook/turbodiffusion.md +++ b/docs/cookbook/turbodiffusion.md @@ -5,7 +5,7 @@ hide: # TurboDiffusion recipes -
    +
    All model families
    diff --git a/docs/cookbook/wan.md b/docs/cookbook/wan.md index c36c325086..7139d02fb0 100644 --- a/docs/cookbook/wan.md +++ b/docs/cookbook/wan.md @@ -5,7 +5,7 @@ hide: # Wan recipes -
    +
    All model families
    diff --git a/docs/cookbook/z-image.md b/docs/cookbook/z-image.md index b53c79ef16..4db21f991c 100644 --- a/docs/cookbook/z-image.md +++ b/docs/cookbook/z-image.md @@ -5,7 +5,7 @@ hide: # Z-Image recipes -
    +
    All model families
    diff --git a/docs/inference/fasth3_rtx_pro_6000.md b/docs/inference/fasth3_rtx_pro_6000.md new file mode 100644 index 0000000000..9debdf0430 --- /dev/null +++ b/docs/inference/fasth3_rtx_pro_6000.md @@ -0,0 +1,260 @@ +# FastH3 NVFP4 on RTX PRO 6000 (sm_120) + +This page covers serving the FastH3 NVFP4 checkpoints on one RTX PRO 6000 +Blackwell (sm_120, 96 GB) with every component resident: the NVFP4 text +encoder, the NVFP4 denoiser and an INT8 light VAE. It also records what was +measured and tried along the way. Every switch is opt-in and defaults to the +existing behavior. + +## Results + +All results are for one RTX PRO 6000 (Modal), a 10.1 s clip at 1344x768 +(243 frames, 73.6k packed tokens), the light INT8 VAE and warm runs. + +| Checkpoint | Denoise | Video decode | End to end | Peak memory | +| --- | ---: | ---: | ---: | ---: | +| [`FastH3-4-step-Preview-v1-VSA-DataFree-NVFP4`](https://huggingface.co/FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree-NVFP4) (4 forwards, VSA 0.9) | 30.1–31.0 s | 9.0 s | **41.2 / 42.4 s** | 79.6 GB | +| [`FastH3-8-Step-V2-NVFP4`](https://huggingface.co/FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4) (8 forwards, VSA 0.8) | 75.6 s | 9.0 s | **86.5 s** | — | + +These end-to-end runs predate the warp-skip kernel change below, which cuts +sparse attention by a further 1.5x, so they are upper bounds. Conditioning +takes about 0.12 s, frame post-processing plus MP4 writing about 0.9 s, and +the first clip at a new shape about 200 s (VAE compile). + +At 480p (124 frames, 15.1k tokens) V2 8-step denoises in 12.5 s on a cold +run, down from 16.8 s warm before this work. + +## Usage + +### 1. Convert the checkpoint + +The published NVFP4 checkpoints use ModelOpt's unified Hugging Face layout. +`convert_minimax_h3_modelopt_nvfp4_dit.py` repacks it into FastVideo's packed +export, `transformer/nvfp4_weights.safetensors`. The calibrated FFN weights +and scales are carried over bit for bit and only the scale bytes are +swizzled. Optionally it also quantizes the BF16 attention projections and the +VSA compression gates: + +```bash +python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \ + --src /models/FastH3-8-Step-V2-NVFP4/transformer \ + --dst /models/fasth3-v2-fv/transformer \ + --quantize-attention --quantize-gate +``` + +Every exported linear is probed through the loader's `mm_fp4` path against a +BF16 matmul with the dequantized weight. Conversion refuses to write when any +error exceeds `--max-probe-error` (default 0.3; the converted V2 and 4-step +checkpoints probe at 0.14). The rest of the model folder (text encoder, +VAEs, schedulers, `fastvideo_inference.json`) is used as-is; the +NVFP4 text encoder comes from +`convert_minimax_h3_text_encoder_nvfp4.py`. + +The flags select the `layer_profile` to load the result with: + +| Flags | Linears in NVFP4 | `layer_profile` | +| --- | --- | --- | +| *(none)* | FFN `fc_in`/`fc_out` | `h3_dit_ffn` | +| `--quantize-attention` | + attention `to_{q,k,v,out}` | `h3_dit` | +| `--quantize-attention --quantize-gate` | + VSA `to_gate_compress` | `h3_dit_vsa` | + +### 2. Generate + +```python +import os +os.environ.update({ + "FASTVIDEO_H3_VSA_FP4": "1", # sparse FP4 attention + "FASTVIDEO_MINIMAX_H3_FUSIONS": "all", # Triton norm/modulate/RoPE/SwiGLU fusions + "FASTVIDEO_NVFP4_MM_BACKEND": "cutlass", # see "FP4 GEMM backend" below + "FASTVIDEO_H3_VAE_TILE_BATCH": "28", # one decoder call per 1344x768 tile grid +}) +from fastvideo import VideoGenerator + +generator = VideoGenerator.from_config({ + "model_path": "/models/fasth3-v2-fv", + "engine": { + "num_gpus": 1, + "quantization": {"transformer_quant": "NVFP4", "layer_profile": "h3_dit_vsa"}, + "compile": {"enabled": False, "vae_enabled": True}, + }, + "pipeline": {"experimental": {"attention_backend": "VIDEO_SPARSE_ATTN_H3", + "VSA_sparsity": 0.8, "VSA_tile_size": 64}}, +}) +generator.generate({"prompt": "...", "sampling": {"height": 768, "width": 1344, "num_frames": 243, + "num_inference_steps": 9}}) +``` + +Use `VSA_sparsity` 0.9 and `num_inference_steps` 5 for the 4-step checkpoint +(see its `fastvideo_inference.json`). Frame counts must be `17n + 5`: 243 +frames is the closest to 10 s. + +## What changed + +### Block-sparse FP4 attention for VSA tiles (`fastvideo-kernel`) + +SageAttention3's sm_120 FP4 kernel (`attn_qat_infer`) gains a block-sparse +forward, `fwd_sparse`, exposed as `sageattn_blackwell_sparse` (head-major +inputs) and `sageattn_blackwell_sparse_bshd` (sequence-major inputs, quantized +in place without a transpose). `vsa_tile_mask_to_fp4_blocks` turns a VSA tile +mask into the kernel's lists: + +- **Block lists.** Query block `m` visits only the 128-token KV blocks in + `q2k_idx[b, h, m, :q2k_num[b, h, m]]`. +- **Quadrant masks for 64-token tiles.** The kernel computes on 128x128 + blocks, but V2 and the 4-step preview use 64-token VSA tiles. Each listed + block carries a 4-bit `q2k_quad` (one bit per 64x64 quadrant). The kernel + masks unselected quadrants to `-inf`, so the result is exactly VSA's tile-64 + semantics. +- **Valid counts per 64-column half.** `kv_valid` gives the valid tokens in + each 64-column half, so partially filled tiles can pad mid-block. +- **Warp-level skipping.** Each MMA warp owns 16 query rows and so sits + inside one 64-row half. A warp skips a listed block that its half did not + select, and the P·V chunk of a key half it did not select. Masked scores + contribute exactly zero, and the warp sharing its tensor-core partition runs + faster meanwhile. This recovers most of the work that pairing 64-token tiles + into 128-token blocks adds. +- **First-visited block.** Lists run in descending block order because the + kernel visits them last entry first. Block 0 (the first prefix tile, which + VSA-H3's exempt mode gives every query) is therefore visited first, so every + row starts from a finite running max. `validate=True` checks this. The model + integration uses only exempt mode. + +The dense and sparse entry points also stop allocating `delta_s`. With Q +smoothing off (the default), each call used to allocate and zero a +`[B, H, L/128, L]` fp32 tensor: 9.5 GB at 73k tokens. Its int32 batch stride +also overflowed the TMA descriptor ("Failed to initialize the TMA descriptor +1", then an illegal instruction), so FP4 attention could not run 10 s 1344x768 +clips at all. A cached `[B, H, 1, L]` zero row read with `per_block_mean=False` +replaces it, and the outputs are bit-identical. + +Correctness (`fastvideo-kernel/tests/test_attn_qat_infer_sparse.py`, RTX PRO +6000): + +| Layout | Error vs token-masked fp32 | Dense FP4 floor | +| --- | ---: | ---: | +| 64-token tiles, odd count, partial tiles | 0.190 | 0.191 | +| 64-token tiles, even count | 0.191 | 0.193 | +| 256-token tiles, partial tile | 0.195 | 0.196 | + +Errors are relative L2. The sparse kernel sits exactly at the FP4 noise +floor; random Gaussian inputs make that floor large. Full block lists +reproduce the dense kernel bit for bit. + +### Model integration (`FASTVIDEO_H3_VSA_FP4=1`) + +`fastvideo/models/dits/minimax_h3_vsa_fp4.py` replaces only the attention core +of `MiniMaxH3Attention`. VSA-H3's tile pooling, top-k mask, exempt prefix and +gated compression branch are unchanged. Per block, it: + +1. Gathers the attention input into tile order once (one `hidden_size`-wide + pass). Pad rows stay zero, so the q/k/v pad rows are exactly zero through + the bias-free projections, RMSNorm and RoPE. +2. Quantizes that input once and shares it between `to_q`, `to_k` and `to_v`. + NVFP4 activations use a unit global scale, so this is exact. +3. Applies QK-norm and RoPE with tile-ordered `cos`/`sin`, computed once per + step. +4. Runs the sparse FP4 kernel on sequence-major tensors and gathers the output + back to packed order before `to_out`. + +This replaces the generic path's concat, four tile scatters and three +transposes. The route applies only to no-grad, non-compiled, single +sequence-parallel-rank calls in exempt mode; everything else keeps the +existing path. + +Two smaller pieces ship alongside it: + +- **Packed gate check.** With `--quantize-gate`, `to_gate_compress` loses its + BF16 weight, so the gate-activity check reads the packed E2M1 bytes instead. +- **Compiled residual.** With `FASTVIDEO_MINIMAX_H3_FUSIONS` enabled, each + block's final `hidden + gate[indices] * ffn_out` runs as one compiled op + instead of materializing the gathered gate. + +### FP4 GEMM backend (`FASTVIDEO_NVFP4_MM_BACKEND`) + +FlashInfer's `mm_fp4(backend="auto")` picks a kernel about 2x slower than +`cutlass` or `cudnn` on sm_120 once activations reach tens of thousands of +rows. At 15k rows all three match. + +| Linear | 73.6k rows: `auto` | 73.6k rows: `cutlass` | 73.6k rows: `cudnn` | 15.1k rows: `auto` | +| --- | ---: | ---: | ---: | ---: | +| `to_q` (5376→7168) | 7.96 ms | 3.96 ms | 4.24 ms | 0.88 ms | +| `to_out` (7168→5376) | 9.07 ms | 4.09 ms | 4.27 ms | 0.92 ms | +| `fc_in` (5376→28672) | 21.54 ms | 16.59 ms | 16.17 ms | 3.01 ms | +| `fc_out` (14336→5376) | 18.06 ms | 8.09 ms | 8.43 ms | 1.65 ms | + +### Batched VAE tile decode (`FASTVIDEO_H3_VAE_TILE_BATCH`) + +The H3 video VAE decodes 256-pixel spatial tiles one at a time. A 1344x768 +clip is a 4x7 grid per temporal chunk, so a 10 s clip is roughly 400 small +decoder calls. The ViT decoder treats batch entries independently, so +`FASTVIDEO_H3_VAE_TILE_BATCH=N` decodes up to `N` equal-shaped tiles per call; +28 covers a full 1344x768 grid. The decoded tiles are the same as per-tile +decoding. + +## Per-block measurements + +One H3 transformer block (hidden 5376, 56 heads, FFN 14336), RTX PRO 6000: + +| Component | 480p, 124 f (15.1k tokens) | 768p, 243 f (73.6k tokens) | +| --- | ---: | ---: | +| VSA Triton BF16 attention (kernel + pooling/mask) | 14.8 ms | 180.1 ms | +| Tile scatter of q/k/v/gate + gather (generic path) | 2.7 ms | 12.7 ms | +| Dense FP4 attention (SageAttention3) | 10.5 ms | 211.9 ms | +| Sparse FP4, quadrant masks | 7.9 ms | 123.4 ms | +| Sparse FP4, quadrant masks + warp skip | — | **83.0 ms** (VSA 0.8) / **49.3 ms** (VSA 0.9) | +| Dense BF16 SDPA | 18.1 ms | — | +| Modulation: eager / fused / compiled | 4.07 / 1.42 / 0.58 ms | 20.3 / 7.0 / 2.8 ms | +| SwiGLU: eager / fused | 1.53 / 0.88 ms | 7.37 / 4.24 ms | +| QK-norm + RoPE: eager / fused | 4.64 / 1.12 ms | 22.4 / 5.2 ms | + +Before this work, a 480p block cost about 40 ms: 17.7 ms of attention, 12 ms +of linears and 10 ms of eager elementwise ops. Over 50 blocks that is 2.0 s +per step, which matches the measured 2.1 s. + +Block density each kernel granularity computes at 768p (fraction of the +dense attention). "Selected" is what VSA needs; the other columns are what +each block shape computes: + +| VSA sparsity | Selected (64x64) | 128x128 blocks | 64-row x 128 | 128 x 64-col | +| --- | ---: | ---: | ---: | ---: | +| 0.8 | 0.222 | 0.434 | 0.317 | 0.313 | +| 0.9 | 0.125 | 0.254 | 0.180 | 0.177 | + +## What was tried and not shipped + +- **Dense FP4 attention for VSA-trained students.** `ATTN_QAT_INFER` does not + build `to_gate_compress`. A VSA-distilled checkpoint such as V2 carries + trained gates, so the strict loader refuses it ("Parameter + ...to_gate_compress.weight not found"). Dense FP4 attention is also slower + than sparse FP4 at 768p (212 vs 83 ms per block). +- **Multi-GPU (Ulysses) FP8 exchange.** On 8x RTX PRO 6000 (PCIe only), NCCL + all-to-all moves about 21 GB/s per GPU, with NCCL P2P on or off. A BF16 + q/k/v/gate exchange at 73.6k tokens therefore costs 24 ms per block, and the + attention output another 6.6 ms; with FP4 payloads q/k/v/gate drop to 7.2 ms. + The branch `h3-sm120-experimental` keeps a sequence-parallel path that: + - sends q/k/v as FP8 with one scale per token and head; + - never sends the VSA gate, applying it on each rank after a small + all-gather of the per-tile compression output; + - returns the attention output as FP8. + + It is estimated at about 18–20 s per 10 s clip for V2 8-step on 8 GPUs. It + has not executed yet (8-GPU capacity was unavailable), so it is not part of + this change. The same branch holds the Modal drivers behind every number on + this page. +- **64-row query blocks.** The kernel traits allow `kBlockM = 64`, which would + remove the query-side pairing waste. Warp-level skipping recovers most of + that waste without a second kernel instantiation, so it was not built. + +## Known limitations + +- **End-to-end quality.** The kernel matches a masked reference at the FP4 + noise floor. Generated videos have not yet been A/B-compared against the + BF16 Triton VSA path on the H3 audio/video metrics. +- **Activation scales.** The packed export drops ModelOpt's calibrated static + `input_scale` and quantizes activations with a unit global scale and dynamic + per-16 block scales, as FastVideo's NVFP4 linears do elsewhere. +- **Decode cost.** Video decode (9 s at 10 s/768p with the light INT8 VAE) is + the next largest cost after denoising. +- **Hardware.** Everything here targets sm_120. The GeForce RTX 5090 shares + the architecture but has 32 GB, which needs a reduced-AdaLN checkpoint and a + non-resident text encoder at this resolution. diff --git a/docs/inference/optimizations.md b/docs/inference/optimizations.md index da98a0313c..adcfc730ae 100644 --- a/docs/inference/optimizations.md +++ b/docs/inference/optimizations.md @@ -198,6 +198,11 @@ The `attn_qat_infer` kernel hard-gates on **sm_120 (consumer Blackwell / RTX 5090)**; on other GPUs the backend logs a notice and falls back to Flash Attention. See the [Attn-QAT paper](https://arxiv.org/abs/2603.00040). +For VSA-distilled MiniMax-H3 students, the same kernel has a block-sparse +forward that runs VSA's 64-token tile selection in FP4 +(`FASTVIDEO_H3_VSA_FP4=1`). See +[FastH3 NVFP4 on RTX PRO 6000](fasth3_rtx_pro_6000.md). + Enable both halves — attention via the env var, linear via `transformer_quant`: ```python diff --git a/examples/inference/basic/basic_compacth3_rtx5090.yaml b/examples/inference/basic/basic_compacth3_rtx5090.yaml new file mode 100644 index 0000000000..46ce251b8f --- /dev/null +++ b/examples/inference/basic/basic_compacth3_rtx5090.yaml @@ -0,0 +1,49 @@ +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_compacth3_rtx5090.yaml +generator: + model_path: ./CompactH3 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: true + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: true + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae +request: + prompt: >- + A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain + peaks above a blue valley filled with thin morning mist. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/compacth3_rtx5090/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml new file mode 100644 index 0000000000..6c92ae7e8e --- /dev/null +++ b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml @@ -0,0 +1,49 @@ +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_compacth3_rtx_pro6000.yaml +generator: + model_path: ./CompactH3 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: true + pipeline: + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae +request: + prompt: >- + A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain + peaks above a blue valley filled with thin morning mist. + negative_prompt: "" + sampling: + seed: 2026 + height: 768 + width: 1344 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/compacth3_rtx_pro6000/ + save_video: true + return_frames: false diff --git a/examples/serving/openai_compacth3_rtx5090.yaml b/examples/serving/openai_compacth3_rtx5090.yaml new file mode 100644 index 0000000000..136fc0af9a --- /dev/null +++ b/examples/serving/openai_compacth3_rtx5090.yaml @@ -0,0 +1,53 @@ +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a \ +# fastvideo serve --config examples/serving/openai_compacth3_rtx5090.yaml +generator: + model_path: ./CompactH3 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: true + image_encoder: true + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: true + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae + +server: + host: 127.0.0.1 + port: 8000 + output_dir: outputs/openai_compacth3_rtx5090 + served_model_name: compacth3 + +default_request: + negative_prompt: "" + sampling: + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + seed: 2026 + output: + return_frames: false diff --git a/examples/serving/openai_compacth3_rtx_pro6000.yaml b/examples/serving/openai_compacth3_rtx_pro6000.yaml new file mode 100644 index 0000000000..0f947feec0 --- /dev/null +++ b/examples/serving/openai_compacth3_rtx_pro6000.yaml @@ -0,0 +1,53 @@ +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a \ +# fastvideo serve --config examples/serving/openai_compacth3_rtx_pro6000.yaml +generator: + model_path: ./CompactH3 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: true + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: true + pipeline: + workload_type: t2v + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae + +server: + host: 127.0.0.1 + port: 8000 + output_dir: outputs/openai_compacth3_rtx_pro6000 + served_model_name: compacth3 + +default_request: + negative_prompt: "" + sampling: + height: 768 + width: 1344 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + seed: 2026 + output: + return_frames: false diff --git a/fastvideo-kernel/attn_qat_infer/api.py b/fastvideo-kernel/attn_qat_infer/api.py index f506d692da..e91196a82f 100644 --- a/fastvideo-kernel/attn_qat_infer/api.py +++ b/fastvideo-kernel/attn_qat_infer/api.py @@ -106,12 +106,32 @@ def pad_to_block_size(x): if enable_smoothing_q: delta_s = torch.matmul(qm, k.transpose(-2, -1)).to(torch.float32).contiguous() else: # used to disable q smoothing - B, H, L, D = q.shape - delta_s = torch.zeros((B, H, L // BLOCK_M, k.shape[2]), device=q.device, dtype=torch.float32) + delta_s = _zero_delta_s(q.shape[0], q.shape[1], k.shape[2], q.device) return q, k, v, delta_s +_ZERO_DELTA_S: dict = {} + + +def _zero_delta_s(batch: int, heads: int, kv_len: int, device: torch.device) -> torch.Tensor: + """Cached all-zero delta_s for unsmoothed Q, read with per_block_mean=False. + + The kernel reads delta_s as a contiguous [B, H, rows, KL] tensor with one + row per query block when per_block_mean is set and a single row otherwise. + With Q smoothing off every row is zero, so one shared row replaces the + [B, H, L/128, KL] tensor that was allocated and zero-filled on every call + (9.5 GB at 73k tokens, whose int32 batch stride also broke the TMA + descriptor). The kernel only reads it. + """ + key = (batch, heads, kv_len, device) + zeros = _ZERO_DELTA_S.get(key) + if zeros is None: + zeros = torch.zeros((batch, heads, 1, kv_len), device=device, dtype=torch.float32) + _ZERO_DELTA_S[key] = zeros + return zeros + + def scale_and_quant_fp4(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: assert x.ndim == 4 B, H, N, D = x.shape @@ -154,6 +174,191 @@ def blockscaled_fp4_attn(qlist: Tuple, softmax_scale, is_causal, per_block_mean, is_bf16, single_level_p_quant) +def blockscaled_fp4_attn_sparse(qlist: Tuple, + klist: Tuple, + vlist: Tuple, + delta_s: torch.Tensor, + KL: int, + q2k_idx: torch.Tensor, + q2k_num: torch.Tensor, + kv_valid: torch.Tensor | None = None, + q2k_quad: torch.Tensor | None = None, + per_block_mean: bool = True, + is_bf16: bool = True, + single_level_p_quant: bool = False, + sm_scale: float | None = None): + softmax_scale = sm_scale if sm_scale is not None else (qlist[0].shape[-1] * 2)**(-0.5) + return fp4attn_cuda.fwd_sparse(qlist[0], klist[0], vlist[0], qlist[1], klist[1], vlist[1], delta_s, KL, None, + softmax_scale, per_block_mean, is_bf16, single_level_p_quant, q2k_idx, q2k_num, + kv_valid, q2k_quad) + + +HALF_N = BLOCK_N // 2 + + +def vsa_tile_mask_to_fp4_blocks( + tile_mask: torch.Tensor, + tile_tokens: int, + tile_valid: torch.Tensor | None = None, + validate: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: + """Convert a VSA tile mask into the FP4 kernel's block lists. + + ``tile_mask`` is a ``[B, H, T, T]`` bool mask over VSA tiles of + ``tile_tokens`` tokens (query tile, key tile); ``tile_tokens`` is 64 or a + multiple of 128. ``tile_valid`` holds each tile's valid token count (valid + tokens first, padding last, as VSA tiling lays them out). The kernel works + on 128x128 blocks, so 64-token tiles are paired: a block is listed when any + of its four 64x64 quadrants is selected, and ``q2k_quad`` carries which. + + Lists run in descending block order because the kernel visits them last + entry first: block 0 (the first prefix tile, which VSA-H3's exempt mode + gives every query) is then visited first, so every query row starts from a + finite running max. ``validate=True`` checks that (host sync). + + Returns ``(q2k_idx [B, H, M, N], q2k_num [B, H, M], kv_valid [2N], + q2k_quad [B, H, M, N] | None)``; the queries cover ``M * 128`` padded rows. + """ + if tile_tokens % HALF_N or (tile_tokens > HALF_N and tile_tokens % BLOCK_N): + raise ValueError(f"tile_tokens={tile_tokens} must be {HALF_N} or a multiple of {BLOCK_N}") + batch, heads, n_tiles, _ = tile_mask.shape + device = tile_mask.device + halves_per_tile = tile_tokens // HALF_N + half_mask = tile_mask + if halves_per_tile > 1: + half_mask = half_mask.repeat_interleave(halves_per_tile, dim=2).repeat_interleave(halves_per_tile, dim=3) + n_halves = n_tiles * halves_per_tile + if tile_valid is None: + half_valid = torch.full((n_halves, ), HALF_N, device=device, dtype=torch.int32) + else: + offsets = torch.arange(halves_per_tile, device=device, dtype=torch.int32) * HALF_N + half_valid = (tile_valid.to(torch.int32)[:, None] - offsets[None, :]).clamp(0, HALF_N).reshape(-1) + if n_halves % 2: + # Pad to whole 128-blocks: the extra key half is empty, the extra query + # half (discarded output) attends block 0 only. + half_mask = F.pad(half_mask, (0, 1, 0, 1), value=False) + half_mask[:, :, -1, 0] = True + half_valid = F.pad(half_valid, (0, 1), value=0) + n_halves += 1 + half_mask = half_mask & (half_valid > 0)[None, None, None, :] + n_blocks = n_halves // 2 + quads = half_mask.view(batch, heads, n_blocks, 2, n_blocks, 2) + weights = torch.tensor([[1, 2], [4, 8]], device=device, dtype=torch.uint8) # [row_half, col_half] + quad = (quads.to(torch.uint8) * weights[None, None, None, :, None, :]).sum(dim=(3, 5), dtype=torch.uint8) + block_mask = quad != 0 + # Descending compaction without a sort: position = running count from the right. + rev = block_mask.flip(-1) + pos = rev.cumsum(-1, dtype=torch.int32) - 1 + q2k_num = (pos[..., -1] + 1).contiguous() + cols = torch.arange(n_blocks - 1, -1, -1, device=device, dtype=torch.int32).expand_as(pos) + slot = torch.where(rev, pos, torch.full_like(pos, n_blocks)).long() + q2k_idx = torch.zeros((batch, heads, n_blocks, n_blocks + 1), device=device, dtype=torch.int32) + q2k_idx.scatter_(-1, slot, cols) + q2k_idx = q2k_idx[..., :n_blocks].contiguous() + kv_valid = half_valid.contiguous() + q2k_quad = None + if halves_per_tile == 1: + q2k_quad = torch.zeros((batch, heads, n_blocks, n_blocks + 1), device=device, dtype=torch.uint8) + q2k_quad.scatter_(-1, slot, quad.flip(-1)) + q2k_quad = q2k_quad[..., :n_blocks].contiguous() + if validate: + if int(q2k_num.min()) < 1: + raise ValueError("every query block must attend to at least one non-empty KV block") + last = (q2k_num - 1).long().unsqueeze(-1) + first_block = q2k_idx.gather(-1, last.int().long()).squeeze(-1).long() + first_quad = (q2k_quad.gather(-1, last).squeeze(-1).int() if q2k_quad is not None else + torch.full_like(first_block, 15, dtype=torch.int32)) + v0 = (kv_valid[2 * first_block] > 0).int() + v1 = (kv_valid[2 * first_block + 1] > 0).int() + row0 = ((first_quad & 1).bool() & v0.bool()) | ((first_quad & 2).bool() & v1.bool()) + row1 = ((first_quad & 4).bool() & v0.bool()) | ((first_quad & 8).bool() & v1.bool()) + if not bool((row0 & row1).all()): + raise ValueError("the first block each query block visits must give both 64-row halves a valid key") + return q2k_idx, q2k_num, kv_valid, q2k_quad + + +def check_sparse_block_lists(q2k_idx: torch.Tensor, q2k_num: torch.Tensor, kv_len: int) -> None: + """Reject block lists the sparse kernel would read out of bounds (host sync). + + Each row needs ``1 <= q2k_num <= q2k_idx.size(-1)`` and every listed index + in ``[0, ceil(kv_len / BLOCK_N))``; a zero count or an out-of-range index + makes the kernel load outside its index row or the KV tensors. + """ + num_kv_blocks = -(-kv_len // BLOCK_N) + if int(q2k_num.min()) < 1 or int(q2k_num.max()) > q2k_idx.size(-1): + raise ValueError(f"q2k_num must be in [1, {q2k_idx.size(-1)}]") + listed = torch.arange(q2k_idx.size(-1), device=q2k_idx.device) < q2k_num.unsqueeze(-1) + idx = q2k_idx[listed] + if idx.numel() and (int(idx.min()) < 0 or int(idx.max()) >= num_kv_blocks): + raise ValueError(f"q2k_idx entries must be in [0, {num_kv_blocks})") + + +def sageattn_blackwell_sparse(q, + k, + v, + q2k_idx: torch.Tensor, + q2k_num: torch.Tensor, + kv_valid: torch.Tensor | None = None, + q2k_quad: torch.Tensor | None = None, + per_block_mean=True, + single_level_p_quant=True, + sm_scale: float | None = None, + validate: bool = True): + """Block-sparse SageAttention3 FP4 forward (non-causal). + + Query block ``m`` (``BLOCK_M`` rows) of each (batch, head) attends only to + the ``BLOCK_N``-token KV blocks in ``q2k_idx[b, h, m, :q2k_num[b, h, m]]``, + restricted to the quadrants in ``q2k_quad`` when given; see + :func:`vsa_tile_mask_to_fp4_blocks`. Q/K/V are ``[B, H, L, D]``. + Block lists are checked with :func:`check_sparse_block_lists` (a host + sync) unless ``validate=False``; pass that only for lists built by + :func:`vsa_tile_mask_to_fp4_blocks`, which are in range by construction. + """ + QL = q.size(2) + KL = k.size(2) + if validate: + check_sparse_block_lists(q2k_idx, q2k_num, KL) + is_bf16 = q.dtype == torch.bfloat16 + q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean) + per_block_mean = delta_s.shape[2] > 1 + qlist_from_cuda = scale_and_quant_fp4(q) + klist_from_cuda = scale_and_quant_fp4_permute(k) + vlist_from_cuda = scale_and_quant_fp4_transpose(v) + o_fp4 = blockscaled_fp4_attn_sparse(qlist_from_cuda, klist_from_cuda, vlist_from_cuda, delta_s, KL, q2k_idx, + q2k_num, kv_valid, q2k_quad, per_block_mean, is_bf16, single_level_p_quant, + sm_scale)[0][:, :, :QL, :].contiguous() + return o_fp4 + + +def sageattn_blackwell_sparse_bshd(q, + k, + v, + q2k_idx: torch.Tensor, + q2k_num: torch.Tensor, + kv_valid: torch.Tensor | None = None, + q2k_quad: torch.Tensor | None = None, + single_level_p_quant=True, + sm_scale: float | None = None, + validate: bool = True) -> torch.Tensor: + """:func:`sageattn_blackwell_sparse` for ``[B, L, H, D]`` inputs, without copies. + + The FP4 quantizers read strided input, so the sequence-major tensors a + linear produces are quantized in place of a transpose + pad. ``L`` must be + a multiple of ``BLOCK_M`` (callers allocate the padding) and Q is + unsmoothed. Returns ``[B, H, L, D]``. + """ + batch, seq_len, heads, _ = q.shape + if seq_len % BLOCK_M: + raise ValueError(f"sequence length {seq_len} must be a multiple of {BLOCK_M}") + if validate: + check_sparse_block_lists(q2k_idx, q2k_num, seq_len) + qh, kh, vh = (x.transpose(1, 2) for x in (q, k, v)) + delta_s = _zero_delta_s(batch, heads, seq_len, q.device) + return blockscaled_fp4_attn_sparse(scale_and_quant_fp4(qh), scale_and_quant_fp4_permute(kh), + scale_and_quant_fp4_transpose(vh), delta_s, seq_len, q2k_idx, q2k_num, kv_valid, + q2k_quad, False, q.dtype == torch.bfloat16, single_level_p_quant, sm_scale)[0] + + def sageattn_blackwell(q, k, v, @@ -191,6 +396,7 @@ def sageattn_blackwell(q, KL = k.size(2) is_bf16 = q.dtype == torch.bfloat16 q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean) + per_block_mean = delta_s.shape[2] > 1 qlist_from_cuda = scale_and_quant_fp4(q) klist_from_cuda = scale_and_quant_fp4_permute(k) vlist_from_cuda = scale_and_quant_fp4_transpose(v) diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/api.cu b/fastvideo-kernel/attn_qat_infer/blackwell/api.cu index 243d876dbf..c7a3624f0b 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/api.cu +++ b/fastvideo-kernel/attn_qat_infer/blackwell/api.cu @@ -204,8 +204,17 @@ void run_mha_fwd(Flash_fwd_params ¶ms, cudaStream_t stream, bool force_split })); } -std::vector -mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // 2) +struct SparseKvLists { + int const *q2k_idx = nullptr; + int const *q2k_num = nullptr; + int q2k_max = 0; + int num_m_blocks = 0; + int const *kv_valid = nullptr; + uint8_t const *q2k_quad = nullptr; +}; + +static std::vector +mha_fwd_impl(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // 2) const at::Tensor &k, // batch_size x seqlen_k x num_heads_k x (head_size // 2) const at::Tensor &v, // batch_size x seqlen_k x num_heads_k x (head_size // 2) const at::Tensor &sfq, @@ -218,7 +227,8 @@ mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size bool is_causal, bool per_block_mean, bool is_bf16, - bool single_level_p_quant=false // If true, use only per-row scale s_P2 (no per-block s_P1) + bool single_level_p_quant, // If true, use only per-row scale s_P2 (no per-block s_P1) + SparseKvLists const &sparse ) { auto dprops = at::cuda::getCurrentDeviceProperties(); @@ -316,6 +326,12 @@ mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // stack-local tensor whose data pointer would dangle after mha_fwd returns // while the async kernel may still be running. params.tile_count_semaphore = nullptr; + params.q2k_idx = sparse.q2k_idx; + params.q2k_num = sparse.q2k_num; + params.q2k_max = sparse.q2k_max; + params.num_m_blocks = sparse.num_m_blocks; + params.kv_valid = sparse.kv_valid; + params.q2k_quad = sparse.q2k_quad; if (seqlen_k > 0) { auto stream = at::cuda::getCurrentCUDAStream().stream(); @@ -341,7 +357,72 @@ mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size +std::vector +mha_fwd(at::Tensor &q, const at::Tensor &k, const at::Tensor &v, + const at::Tensor &sfq, const at::Tensor &sfk, const at::Tensor &sfv, + const at::Tensor &delta_s, int unpadded_k, c10::optional &out_, + const float softmax_scale, bool is_causal, bool per_block_mean, bool is_bf16, + bool single_level_p_quant=false) { + return mha_fwd_impl(q, k, v, sfq, sfk, sfv, delta_s, unpadded_k, out_, softmax_scale, + is_causal, per_block_mean, is_bf16, single_level_p_quant, SparseKvLists{}); +} + +// Block-sparse forward: query block m of (b, h) attends only to the KV blocks +// listed in q2k_idx[b, h, m, :q2k_num[b, h, m]] (BLOCK_M x BLOCK_N granularity). +// kv_valid, if given, holds the valid token count of each 64-column half of +// every KV block ([2 * num_kv_blocks], valid tokens first within a half). +// q2k_quad, if given (uint8, same shape as q2k_idx), restricts each listed +// block to the 64x64 quadrants whose bit (2 * row_half + col_half) is set, so +// 64-token VSA tiles run on 128x128 blocks. The block a list visits first +// (its last entry) must leave every query row at least one valid key. +// Non-causal only. +std::vector +mha_fwd_sparse(at::Tensor &q, const at::Tensor &k, const at::Tensor &v, + const at::Tensor &sfq, const at::Tensor &sfk, const at::Tensor &sfv, + const at::Tensor &delta_s, int unpadded_k, c10::optional &out_, + const float softmax_scale, bool per_block_mean, bool is_bf16, + bool single_level_p_quant, + const at::Tensor &q2k_idx, const at::Tensor &q2k_num, + c10::optional &kv_valid_, + c10::optional &q2k_quad_) { + const int batch_size = q.size(0); + const int num_heads = q.size(1); + const int num_m_blocks = (q.size(2) + flash::BLOCK_M - 1) / flash::BLOCK_M; + const int num_n_blocks = (k.size(2) + flash::BLOCK_N - 1) / flash::BLOCK_N; + for (auto const *t : {&q2k_idx, &q2k_num}) { + TORCH_CHECK(t->scalar_type() == torch::kInt32, "q2k_idx / q2k_num must be int32"); + CHECK_DEVICE((*t)); CHECK_CONTIGUOUS((*t)); + } + TORCH_CHECK(q2k_idx.dim() == 4, "q2k_idx must be [batch, heads, num_m_blocks, max_kv_blocks]"); + TORCH_CHECK(q2k_idx.size(0) == batch_size && q2k_idx.size(1) == num_heads && q2k_idx.size(2) == num_m_blocks, + "q2k_idx leading dims must be [batch, heads, ceil(seqlen_q / BLOCK_M)]"); + TORCH_CHECK(q2k_idx.size(3) >= 1 && q2k_idx.size(3) <= num_n_blocks, "q2k_idx last dim must be in [1, num_kv_blocks]"); + CHECK_SHAPE(q2k_num, batch_size, num_heads, num_m_blocks); + SparseKvLists sparse; + sparse.q2k_idx = q2k_idx.data_ptr(); + sparse.q2k_num = q2k_num.data_ptr(); + sparse.q2k_max = q2k_idx.size(3); + sparse.num_m_blocks = num_m_blocks; + if (kv_valid_.has_value()) { + auto const &kv_valid = kv_valid_.value(); + TORCH_CHECK(kv_valid.scalar_type() == torch::kInt32, "kv_valid must be int32"); + CHECK_DEVICE(kv_valid); CHECK_CONTIGUOUS(kv_valid); + CHECK_SHAPE(kv_valid, 2 * num_n_blocks); + sparse.kv_valid = kv_valid.data_ptr(); + } + if (q2k_quad_.has_value()) { + auto const &q2k_quad = q2k_quad_.value(); + TORCH_CHECK(q2k_quad.scalar_type() == torch::kUInt8, "q2k_quad must be uint8"); + CHECK_DEVICE(q2k_quad); CHECK_CONTIGUOUS(q2k_quad); + TORCH_CHECK(q2k_quad.sizes() == q2k_idx.sizes(), "q2k_quad must match q2k_idx's shape"); + sparse.q2k_quad = q2k_quad.data_ptr(); + } + return mha_fwd_impl(q, k, v, sfq, sfk, sfv, delta_s, unpadded_k, out_, softmax_scale, + /*is_causal=*/false, per_block_mean, is_bf16, single_level_p_quant, sparse); +} + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.doc() = "FlashAttention"; m.def("fwd", &mha_fwd, "Forward pass"); + m.def("fwd_sparse", &mha_fwd_sparse, "Block-sparse forward pass (non-causal)"); } \ No newline at end of file diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h b/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h index 675608726a..d237499c34 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h @@ -185,14 +185,14 @@ __global__ void __launch_bounds__(Ktraits::kNWarps * cutlass::NumThreadsPerWarp, auto block_coord = work_tile_info.get_block_coord(scheduler_params); auto [m_block, bidh, bidb] = block_coord; - int n_block_max = collective_mainloop.get_n_block_max(mainloop_params, m_block); - if (Is_causal && n_block_max <= 0) { // We exit early and write 0 to gO and -inf to gLSE. + int n_block_count = collective_mainloop.get_n_block_count(mainloop_params, m_block, bidh, bidb); + if (Is_causal && n_block_count <= 0) { // We exit early and write 0 to gO and -inf to gLSE. collective_epilogue.store_zero(epilogue_params, threadIdx.x - NumCopyThreads, block_coord); continue; } collective_mainloop.mma(mainloop_params, pipeline_q, pipeline_k, pipeline_v, smem_pipe_read_q, smem_pipe_read_k, smem_pipe_read_v, - tOrO, softmax_fused, n_block_max, threadIdx.x - NumCopyThreads, work_idx, m_block, shared_storage); + tOrO, softmax_fused, n_block_count, threadIdx.x - NumCopyThreads, work_idx, m_block, bidh, bidb, shared_storage); barrier_o.wait(); collective_epilogue.mma_store(shared_storage, tiled_mma_pv, tOrO, threadIdx.x - NumCopyThreads); barrier_o.arrive(); diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/launch.h b/fastvideo-kernel/attn_qat_infer/blackwell/launch.h index 91398c5f44..edc167afca 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/launch.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/launch.h @@ -62,7 +62,9 @@ void run_flash_fwd(Flash_fwd_params ¶ms, cudaStream_t stream) { static_cast(params.delta_s_ptr), {params.seqlen_s, params.seqlen_k, params.h_k, params.b}, {params.ds_row_stride, _1{}, params.ds_head_stride, params.ds_batch_stride}, - params.scale_softmax_log2 + params.scale_softmax_log2, + params.q2k_idx, params.q2k_num, params.q2k_max, + params.num_m_blocks, params.h, params.kv_valid, params.q2k_quad }); typename CollectiveEpilogue::Params epilogue_params = CollectiveEpilogue::to_underlying_arguments({ diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h b/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h index fa2240422c..51d9ff23a3 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h @@ -174,6 +174,13 @@ struct CollectiveMainloopFwd { ShapeQKV const shape_ds; StrideQKV const stride_ds; float const softmax_scale_log2; + int const* ptr_q2k_idx{nullptr}; + int const* ptr_q2k_num{nullptr}; + int q2k_max{0}; + int num_m_blocks{0}; + int num_heads{0}; + int const* ptr_kv_valid{nullptr}; + uint8_t const* ptr_q2k_quad{nullptr}; }; // Device side kernel params @@ -194,6 +201,13 @@ struct CollectiveMainloopFwd { TMA_SFVt tma_load_SFVt; TMA_DS tma_load_DS; float const softmax_scale_log2; + int const* ptr_q2k_idx; + int const* ptr_q2k_num; + int q2k_max; + int num_m_blocks; + int num_heads; + int const* ptr_kv_valid; + uint8_t const* ptr_q2k_quad; }; @@ -261,7 +275,9 @@ struct CollectiveMainloopFwd { tma_load_K, tma_load_sfk, tma_load_Vt, tma_load_sfvt, tma_load_ds, - args.softmax_scale_log2}; + args.softmax_scale_log2, + args.ptr_q2k_idx, args.ptr_q2k_num, args.q2k_max, + args.num_m_blocks, args.num_heads, args.ptr_kv_valid, args.ptr_q2k_quad}; } /// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance @@ -290,6 +306,51 @@ struct CollectiveMainloopFwd { return n_block_max; } + // Number of KV blocks query block m_block visits: all of them when dense, + // its index-list length when block-sparse. + CUTLASS_DEVICE + int get_n_block_count(Params const& mainloop_params, int m_block, int bidh, int bidb) { + if (mainloop_params.ptr_q2k_idx == nullptr) { + return get_n_block_max(mainloop_params, m_block); + } + return mainloop_params.ptr_q2k_num[(bidb * mainloop_params.num_heads + bidh) * mainloop_params.num_m_blocks + m_block]; + } + + // KV block visited at iteration i (iterations run from count-1 down to 0). + CUTLASS_DEVICE + int get_kv_block(Params const& mainloop_params, int m_block, int bidh, int bidb, int i) { + if (mainloop_params.ptr_q2k_idx == nullptr) { + return i; + } + int64_t const row = (int64_t(bidb) * mainloop_params.num_heads + bidh) * mainloop_params.num_m_blocks + m_block; + return mainloop_params.ptr_q2k_idx[row * mainloop_params.q2k_max + i]; + } + + // Valid key columns of each 64-column half of KV block n_block (valid tokens + // first within a half). kv_valid, when given, stores two counts per block so + // 64-token VSA tiles can pad mid-block; otherwise only the sequence tail is + // masked. + CUTLASS_DEVICE + int get_kv_valid_half(Params const& mainloop_params, int n_block, int half, int unpadded_seqlen_k) { + static constexpr int kBlockN = get<1>(TileShape_MNK{}); + static constexpr int kHalfN = kBlockN / 2; + if (mainloop_params.ptr_kv_valid != nullptr) { + return mainloop_params.ptr_kv_valid[2 * n_block + half]; + } + return max(0, min(kHalfN, unpadded_seqlen_k - n_block * kBlockN - half * kHalfN)); + } + + // Quadrant mask of list entry i: bit (2 * row_half + col_half) is set when the + // 64-row query half attends the 64-column key half. 0xF when absent. + CUTLASS_DEVICE + int get_quad(Params const& mainloop_params, int m_block, int bidh, int bidb, int i) { + if (mainloop_params.ptr_q2k_quad == nullptr) { + return 0xF; + } + int64_t const row = (int64_t(bidb) * mainloop_params.num_heads + bidh) * mainloop_params.num_m_blocks + m_block; + return mainloop_params.ptr_q2k_quad[row * mainloop_params.q2k_max + i]; + } + template CUTE_HOST_DEVICE constexpr auto @@ -450,7 +511,7 @@ struct CollectiveMainloopFwd { auto [m_block, bidh, bidb] = work_tile_info.get_block_coord(scheduler_params); - int n_block_max = get_n_block_max(mainloop_params, m_block); + int n_block_count = get_n_block_count(mainloop_params, m_block, bidh, bidb); Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.begin()), SmemLayoutQ{}); Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.begin()), SmemLayoutK{}); @@ -506,7 +567,8 @@ struct CollectiveMainloopFwd { Tensor tDSsDS = group_modes<0, 3>(block_tma_ds.partition_D(sDS)); uint16_t mcast_mask_kv = 0; - int n_block = n_block_max - 1; + int n_iter = n_block_count - 1; + int n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); int lane_predicate = cute::elect_one_sync(); if (lane_predicate) { pipeline_q.producer_acquire(smem_pipe_write_q); @@ -529,11 +591,12 @@ struct CollectiveMainloopFwd { ++smem_pipe_write_v; } - n_block--; + --n_iter; if (lane_predicate) { // CUTLASS_PRAGMA_NO_UNROLL #pragma unroll 2 - for (; n_block >= 0; --n_block) { + for (; n_iter >= 0; --n_iter) { + n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); pipeline_k.producer_acquire(smem_pipe_write_k); copy(mainloop_params.tma_load_K.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv), tKgK(_, n_block), tKsK(_, smem_pipe_write_k.index())); @@ -585,6 +648,8 @@ struct CollectiveMainloopFwd { int thread_idx, int work_idx, int m_block, + int bidh, + int bidb, SharedStorage& shared_storage ) { @@ -667,7 +732,21 @@ struct CollectiveMainloopFwd { int const seqlen_q = get<0>(mainloop_params.shape_Q); int const seqlen_k = get<0>(mainloop_params.shape_K); int const unpadded_seqlen_k = get<0>(mainloop_params.unpadded_shape_K); - int n_block = n_block_count - 1; + int n_iter = n_block_count - 1; + int n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); + bool const per_block_masking = mainloop_params.ptr_q2k_idx != nullptr || mainloop_params.ptr_kv_valid != nullptr; + static_assert(kBlockM == 128 && kBlockN == 128, "quadrant masking assumes 128x128 blocks"); + // Each MMA warp owns 16 consecutive query rows, so a warp lies in one + // 64-row half. With quadrant lists a warp skips blocks its half did not + // select and the P.V chunk of a key half it did not select: masked + // scores contribute exactly zero, and the warp sharing its tensor-core + // partition runs faster meanwhile. + int const my_row_half = __shfl_sync(0xffffffff, [&] { + Tensor cS0 = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{})); + Tensor tScS0 = thread_mma_qk.partition_C(cS0); + return int(get<0>(tScS0(0))) >= kBlockM / 2 ? 1 : 0; + }(), 0); + bool const quad_skip = mainloop_params.ptr_q2k_quad != nullptr; auto copy_k_block = [&](auto block_id) { auto tSsK_stage = tSsK(_, _, _, smem_pipe_read_k.index()); @@ -744,6 +823,26 @@ struct CollectiveMainloopFwd { int local = c & 31; return (c & ~31) | (local & 1) | ((local & 24) >> 2) | ((local & 6) << 2); }; + // Sparse lists may visit a block whose 64-token halves are partially valid + // (VSA tile tails, short prefix chunks) or that only one query half + // selected (64-token tiles), at any iteration. + auto apply_sparse_mask = [&](auto& acc, int n_blk, int it) { + int const quad = get_quad(mainloop_params, m_block, bidh, bidb, it); + int const valid0 = get_kv_valid_half(mainloop_params, n_blk, 0, unpadded_seqlen_k); + int const valid1 = get_kv_valid_half(mainloop_params, n_blk, 1, unpadded_seqlen_k); + if (quad == 0xF && valid0 == kBlockN / 2 && valid1 == kBlockN / 2) { return; } + Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{})); + Tensor tScS = thread_mma_qk.partition_C(cS); + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < size(acc); ++i) { + int const col = actual_col(int(get<1>(tScS(i)))); + int const col_half = col >= kBlockN / 2; + int const row_half = int(get<0>(tScS(i))) >= kBlockM / 2; + bool const keep = ((quad >> (2 * row_half + col_half)) & 1) + && (col & (kBlockN / 2 - 1)) < (col_half ? valid1 : valid0); + if (!keep) { acc(i) = -INFINITY; } + } + }; { Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{})); Tensor tScS = thread_mma_qk.partition_C(cS); @@ -751,7 +850,9 @@ struct CollectiveMainloopFwd { for (int i = 0; i < size(tSrS); ++i) { int col = actual_col(int(get<1>(tScS(i)))); if constexpr (!Is_causal) { // Just masking based on col - if (col >= int(unpadded_seqlen_k - n_block * kBlockN)) { tSrS(i) = -INFINITY; } + if (!per_block_masking) { + if (col >= int(unpadded_seqlen_k - n_block * kBlockN)) { tSrS(i) = -INFINITY; } + } } else { if (col >= std::min(seqlen_k - n_block * kBlockN, col_limit_causal(int(get<0>(tScS(i))), n_block))) { @@ -760,6 +861,9 @@ struct CollectiveMainloopFwd { } } } + if constexpr (!Is_causal) { + if (per_block_masking) { apply_sparse_mask(tSrS, n_block, n_iter); } + } auto quantize = [&](auto mma_k, auto acc_conversion_view) { Tensor AbsMaxP_stagek = AbsMaxP(_, make_coord(_, _, mma_k)); Tensor acc_conversion_stagek = acc_conversion_view(_, _, mma_k); @@ -827,11 +931,12 @@ struct CollectiveMainloopFwd { } } - n_block--; + --n_iter; constexpr int n_masking_steps = !Is_causal ? 1 : cute::ceil_div(kBlockM, kBlockN) + 1; // // Only go through these if Is_causal, since n_masking_steps = 1 when !Is_causal CUTLASS_PRAGMA_UNROLL - for (int masking_step = 0; masking_step < n_masking_steps - 1 && n_block >= 0; ++masking_step, --n_block) { + for (int masking_step = 0; masking_step < n_masking_steps - 1 && n_iter >= 0; ++masking_step, --n_iter) { + n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{})); Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout())); consumer_wait(pipeline_k, smem_pipe_read_k); @@ -878,7 +983,22 @@ struct CollectiveMainloopFwd { } #pragma unroll 1 - for (; n_block >= 0; --n_block) { + for (; n_iter >= 0; --n_iter) { + n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); + int const row_bits = quad_skip + ? (get_quad(mainloop_params, m_block, bidh, bidb, n_iter) >> (2 * my_row_half)) & 3 + : 3; + if (row_bits == 0) { + // Neither key half is selected for this warp's rows: its softmax + // state and output are unchanged. Keep the pipeline in step. + consumer_wait(pipeline_k, smem_pipe_read_k); + pipeline_k.consumer_release(smem_pipe_read_k); + ++smem_pipe_read_k; + consumer_wait(pipeline_v, smem_pipe_read_v); + pipeline_v.consumer_release(smem_pipe_read_v); + ++smem_pipe_read_v; + continue; + } Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{})); Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout())); consumer_wait(pipeline_k, smem_pipe_read_k); @@ -896,16 +1016,21 @@ struct CollectiveMainloopFwd { } } + if (per_block_masking) { apply_sparse_mask(tSrS, n_block, n_iter); } softmax_fused.template online_softmax_with_quant(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2); Tensor tOrO = make_fragment_like(tOrO_store); + clear(tOrO); consumer_wait(pipeline_v, smem_pipe_read_v); copy_v_block(_0{}); quantize(_0{}, tSrS_converion_view); CUTLASS_PRAGMA_UNROLL for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) { - cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)), - make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO); + // v_block spans one 64-column key half (P's K mode is 2 x 64). + if ((row_bits >> v_block) & 1) { + cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)), + make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO); + } if (v_block < size<2>(tOrP) - 1) { copy_v_block(v_block + 1); quantize(v_block + 1, tSrS_converion_view); diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/params.h b/fastvideo-kernel/attn_qat_infer/blackwell/params.h index fcf9bf8981..867aa7b257 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/params.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/params.h @@ -114,6 +114,23 @@ int * __restrict__ seqused_k; int *__restrict__ blockmask; + + // Block-sparse KV iteration (non-causal only). When q2k_idx is null the + // kernel is dense. Otherwise query block m of (batch b, head h) visits + // the q2k_num[(b*h + h)*num_m_blocks + m] KV blocks listed at + // q2k_idx[((b*h + h)*num_m_blocks + m)*q2k_max + i], in any order. + int const *__restrict__ q2k_idx; + int const *__restrict__ q2k_num; + int q2k_max; + int num_m_blocks; + // Optional valid token count of each 64-column half of every KV block, + // kv_valid[2*n + half] (valid tokens first within a half); null means only + // the sequence tail beyond the unpadded key length is masked. + int const *__restrict__ kv_valid; + // Optional quadrant mask per list entry (same layout as q2k_idx): bit + // (2*row_half + col_half) set when that 64-row query half attends that + // 64-column key half. Lets 64-token VSA tiles run on 128x128 blocks. + uint8_t const *__restrict__ q2k_quad; // The K_new and V_new matrices. void * __restrict__ knew_ptr; diff --git a/fastvideo-kernel/tests/test_attn_qat_infer_sparse.py b/fastvideo-kernel/tests/test_attn_qat_infer_sparse.py new file mode 100644 index 0000000000..3684b2156b --- /dev/null +++ b/fastvideo-kernel/tests/test_attn_qat_infer_sparse.py @@ -0,0 +1,155 @@ +#!/usr/bin/env python3 +"""Correctness tests for the block-sparse SageAttention3 FP4 forward (``fwd_sparse``). + +Covers the VSA tile layouts MiniMax-H3 uses: 64-token tiles carried by 64x64 +quadrant masks on the kernel's 128x128 blocks (odd and even tile counts), +256-token tiles, partially valid tiles, and an exempt prefix that every query +attends. The sparse kernel is compared with a token-masked fp32 reference and +must stay at the dense FP4 kernel's own error on the same data. Also checks +that full block lists reproduce the dense kernel bit for bit and that the +sequence-major entry point matches the head-major one bit for bit. + +Requires a Blackwell GPU (sm_120a) and the fp4attn_cuda / fp4quant_cuda +extensions built via ``cd fastvideo-kernel && ./build.sh``. + + pytest tests/test_attn_qat_infer_sparse.py -v +""" + +import math +import os +import sys + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +import pytest +import torch +import torch.nn.functional as F + +pytest.importorskip("fp4attn_cuda", reason="ATTN_QAT_INFER FP4 kernels require a sm_120a build") + +from attn_qat_infer.api import (BLOCK_M, BLOCK_N, check_sparse_block_lists, sageattn_blackwell, + sageattn_blackwell_sparse, sageattn_blackwell_sparse_bshd, + vsa_tile_mask_to_fp4_blocks) + +DEVICE = torch.device("cuda") +HEAD_DIM = 128 + + +def _tile_sizes(prefix: tuple[int, ...], video_tiles: int, tile: int, partial: dict[int, int]) -> torch.Tensor: + """Valid tokens per tile: segment-pure prefix chunks, then video tiles (some partial).""" + sizes = [] + for segment in prefix: + full, rem = divmod(segment, tile) + sizes += [tile] * full + ([rem] if rem else []) + sizes += [partial.get(index, tile) for index in range(video_tiles)] + return torch.tensor(sizes, dtype=torch.int32, device=DEVICE) + + +def _exempt_mask(n_prefix: int, n_tiles: int, sparsity: float, heads: int, gen: torch.Generator) -> torch.Tensor: + """VSA-H3 exempt selection: prefix rows dense, prefix keys always, top-k video keys.""" + n_video = n_tiles - n_prefix + k_video = max(1, math.ceil((1 - sparsity) * n_video)) + scores = torch.rand(1, heads, n_tiles, n_video, generator=gen, device=DEVICE) + mask = torch.zeros(1, heads, n_tiles, n_tiles, dtype=torch.bool, device=DEVICE) + mask.scatter_(-1, scores.topk(k_video, dim=-1).indices + n_prefix, True) + mask[..., :n_prefix] = True + mask[:, :, :n_prefix, :] = True + return mask + + +def _case(tile: int, prefix: tuple[int, ...], video_tiles: int, partial: dict[int, int], heads: int = 4): + gen = torch.Generator(device=DEVICE).manual_seed(tile * 1000 + video_tiles) + sizes = _tile_sizes(prefix, video_tiles, tile, partial) + n_tiles = sizes.numel() + n_prefix = n_tiles - video_tiles + rows = math.ceil(n_tiles * tile / BLOCK_M) * BLOCK_M + valid = torch.zeros(rows, dtype=torch.bool, device=DEVICE) + valid[:n_tiles * tile] = (torch.arange(tile, device=DEVICE)[None, :] < sizes[:, None]).reshape(-1) + q, k, v = (torch.randn(1, heads, rows, HEAD_DIM, device=DEVICE, dtype=torch.bfloat16, generator=gen) + for _ in range(3)) + for x in (q, k, v): + x[:, :, ~valid] = 0 + mask = _exempt_mask(n_prefix, n_tiles, 0.8, heads, gen) + return q, k, v, sizes, mask, valid + + +def _rel_l2(a: torch.Tensor, b: torch.Tensor) -> float: + return ((a.float() - b.float()).norm() / b.float().norm()).item() + + +@pytest.mark.parametrize( + ("tile", "prefix", "video_tiles", "partial"), + [ + (64, (77, 46), 19, { + 3: 40, + 11: 1 + }), # odd tile count, short prefix chunks + (64, (64, 40), 9, {}), # even tile count + (256, (100, 70), 4, { + 1: 200 + }), # 256-token tiles, partial tile + ], +) +def test_sparse_matches_masked_reference_at_fp4_floor(tile, prefix, video_tiles, partial) -> None: + q, k, v, sizes, mask, valid = _case(tile, prefix, video_tiles, partial) + q2k_idx, q2k_num, kv_valid, q2k_quad = vsa_tile_mask_to_fp4_blocks(mask, tile, sizes, validate=True) + assert (q2k_quad is not None) == (tile < BLOCK_N) + out = sageattn_blackwell_sparse(q.clone(), k.clone(), v.clone(), q2k_idx, q2k_num, kv_valid, q2k_quad) + + n_tok = sizes.numel() * tile + token_tile = torch.arange(sizes.numel(), device=DEVICE).repeat_interleave(tile) + token_mask = mask[:, :, token_tile][:, :, :, token_tile] & valid[None, None, None, :n_tok] + ref = F.scaled_dot_product_attention(q[:, :, :n_tok].float(), + k[:, :, :n_tok].float(), + v[:, :, :n_tok].float(), + attn_mask=token_mask) + rows = valid[:n_tok] + assert torch.isfinite(out[:, :, :n_tok][:, :, rows]).all() + sparse_err = _rel_l2(out[:, :, :n_tok][:, :, rows], ref[:, :, rows]) + + # FP4 floor: the dense kernel against dense fp32 attention on the valid tokens. + qv, kv, vv = (x[:, :, valid].contiguous() for x in (q, k, v)) + floor = _rel_l2(sageattn_blackwell(qv, kv, vv), F.scaled_dot_product_attention(qv.float(), kv.float(), vv.float())) + assert sparse_err <= 1.05 * floor + 1e-3, (sparse_err, floor) + + +def test_full_lists_reproduce_dense_bitwise() -> None: + gen = torch.Generator(device=DEVICE).manual_seed(0) + q, k, v = (torch.randn(1, 4, 2048, HEAD_DIM, device=DEVICE, dtype=torch.bfloat16, generator=gen) for _ in range(3)) + n_blocks = 2048 // BLOCK_N + q2k_idx = torch.arange(n_blocks, device=DEVICE, dtype=torch.int32).expand(1, 4, 2048 // BLOCK_M, + n_blocks).contiguous() + q2k_num = torch.full((1, 4, 2048 // BLOCK_M), n_blocks, device=DEVICE, dtype=torch.int32) + dense = sageattn_blackwell(q.clone(), k.clone(), v.clone()) + sparse = sageattn_blackwell_sparse(q.clone(), k.clone(), v.clone(), q2k_idx, q2k_num) + assert torch.equal(dense, sparse) + + +def test_bshd_entry_matches_bhsd() -> None: + q, k, v, sizes, mask, _ = _case(64, (77, 46), 19, {3: 40}) + lists = vsa_tile_mask_to_fp4_blocks(mask, 64, sizes) + bhsd = sageattn_blackwell_sparse(q.clone(), k.clone(), v.clone(), *lists) + bshd = sageattn_blackwell_sparse_bshd(*(x.transpose(1, 2).contiguous() for x in (q, k, v)), *lists) + assert torch.equal(bhsd, bshd) + + +def test_first_visited_block_check_rejects_unanchored_rows() -> None: + # Without the exempt prefix a query block may start from a block that masks + # one of its 64-row halves entirely; validate=True must refuse it. + mask = torch.zeros(1, 1, 4, 4, dtype=torch.bool, device=DEVICE) + mask[0, 0, 0, 3] = True # query tile 0 -> key tile 3 only + mask[0, 0, 1, 1] = True # query tile 1 -> key tile 1 only (block 0, other quadrant) + mask[0, 0, 2:, 2:] = True + with pytest.raises(ValueError, match="first block"): + vsa_tile_mask_to_fp4_blocks(mask, 64, torch.full((4, ), 64, dtype=torch.int32, device=DEVICE), validate=True) + + +def test_block_list_check_rejects_out_of_bounds_lists() -> None: + idx = torch.zeros((1, 1, 2, 2), dtype=torch.int32, device=DEVICE) + num = torch.ones((1, 1, 2), dtype=torch.int32, device=DEVICE) + idx[..., 1] = 99 # beyond q2k_num: never read, so not checked + check_sparse_block_lists(idx, num, 2 * BLOCK_N) + with pytest.raises(ValueError, match="q2k_num"): + check_sparse_block_lists(idx, torch.zeros_like(num), 2 * BLOCK_N) + with pytest.raises(ValueError, match="q2k_idx"): + check_sparse_block_lists(idx, num + 1, 2 * BLOCK_N) diff --git a/fastvideo/api/compat.py b/fastvideo/api/compat.py index 07d4762cbf..32c13ccedb 100644 --- a/fastvideo/api/compat.py +++ b/fastvideo/api/compat.py @@ -292,8 +292,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An # typed surface accepts a string and does the wiring here so # downstream code can rely on a single source of truth. from fastvideo.layers.quantization import get_quantization_config + from fastvideo.layers.quantization.nvfp4_config import NVFP4Config _resolved_quant_cls = get_quantization_config(quantization.transformer_quant) - kwargs["transformer_quant"] = _resolved_quant_cls() + quant_instance = _resolved_quant_cls() + if quantization.layer_profile is not None: + if not isinstance(quant_instance, NVFP4Config): + raise ValueError("engine.quantization.layer_profile is only valid with transformer_quant NVFP4, " + f"got {type(quant_instance).__name__}") + quant_instance = NVFP4Config( + layer_profile=quantization.layer_profile, + retain_original_weights=quant_instance.retain_original_weights, + ) + kwargs["transformer_quant"] = quant_instance components = normalized.pipeline.components if components.pipeline_config_path is not None: diff --git a/fastvideo/api/schema.py b/fastvideo/api/schema.py index 77224d0368..521039785f 100644 --- a/fastvideo/api/schema.py +++ b/fastvideo/api/schema.py @@ -76,6 +76,11 @@ class CompileConfig: class QuantizationConfig: text_encoder_quant: str | None = None transformer_quant: str | None = None + # Forwarded onto ``NVFP4Config`` when ``transformer_quant`` is ``NVFP4``. + # ``h3_dit`` selects the packed MiniMax-H3 attention+FFN export. + # ``h3_dit_ffn`` selects a packed FFN-only export (attention stays dense). + # ``h3_dit_vsa`` is ``h3_dit`` plus the VSA compression gates. + layer_profile: str | None = None @dataclass diff --git a/fastvideo/attention/backends/minimax_h3_sparse_int8.py b/fastvideo/attention/backends/minimax_h3_sparse_int8.py new file mode 100644 index 0000000000..5854d25196 --- /dev/null +++ b/fastvideo/attention/backends/minimax_h3_sparse_int8.py @@ -0,0 +1,189 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Experimental sm89 tile-64 VSA with BF16 or INT8 QK and FP32 accumulation. + +Retains each query tile's original key selection and masks partial key tiles. +Unlike a 128-query adapter, it adds no attention blocks. Q/K use per-token +scales; K centering is a softmax-invariant shift. V uses one scale per head +and channel, so its dequantization can be applied once in the epilogue. +BF16 PV is the default: FP8 PV had excessive error on real H3 inputs. +Numerical validation and same-seed clip review are required before enabling. +""" +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _quantize_qk(X, Mean, VBS, Y, Scale, L: tl.constexpr, D: tl.constexpr, H: tl.constexpr, XB: tl.constexpr, + XH: tl.constexpr, XS: tl.constexpr, XD: tl.constexpr, CENTER: tl.constexpr, ROWS: tl.constexpr): + hz = tl.program_id(1) + rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS) + cols = tl.arange(0, D) + offset = (hz // H) * XB + (hz % H) * XH + rows[:, None] * XS + cols[None, :] * XD + x = tl.load(X + offset, rows[:, None] < L, 0).to(tl.float32) + if CENTER: + mean = tl.load(Mean + hz * D + cols) + valid_size = tl.load(VBS + rows // 64, rows < L, 0) + x = tl.where((rows % 64 < valid_size)[:, None], x - mean[None, :], 0.0) + scale = tl.maximum(tl.max(tl.abs(x), 1) / 127.0, 1e-8) + y = tl.floor(x / scale[:, None] + 0.5).to(tl.int8) + tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], y, rows[:, None] < L) + tl.store(Scale + hz * L + rows, scale, rows < L) + + +@triton.jit +def _quantize_v(X, Scale, Y, L: tl.constexpr, D: tl.constexpr, ROWS: tl.constexpr): + hz = tl.program_id(1) + rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS) + cols = tl.arange(0, D) + scale = tl.load(Scale + hz * D + cols) + x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :], rows[:, None] < L, 0).to(tl.float32) + tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], (x / scale[None, :]).to(tl.float8e4nv), rows[:, None] + < L) + + +@triton.jit +def _quantize_v_tiles(X, Y, Scale, L: tl.constexpr, D: tl.constexpr): + tile, hz = tl.program_id(0), tl.program_id(1) + rows = tile * 64 + tl.arange(0, 64) + cols = tl.arange(0, D) + x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :]).to(tl.float32) + scale = tl.maximum(tl.max(tl.abs(x), 0) / 448.0, 1e-8) + tl.store(Scale + (hz * (L // 64) + tile) * D + cols, scale) + tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], (x / scale[None, :]).to(tl.float8e4nv)) + + +@triton.autotune( + configs=[triton.Config({}, num_warps=w, num_stages=s) for w, s in ((4, 2), (4, 3), (4, 4), (8, 2), (8, 3))], + key=["L", "D"]) +@triton.jit +def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr, H: tl.constexpr, + VB: tl.constexpr, VH: tl.constexpr, VS_ROW: tl.constexpr, VD: tl.constexpr, INT8_QK: tl.constexpr, + FP8_PV: tl.constexpr, V_TILE: tl.constexpr, P_DYNAMIC: tl.constexpr): + tile, hz = tl.program_id(0), tl.program_id(1) + nt: tl.constexpr = L // 64 + rows = tile * 64 + tl.arange(0, 64) + cols = tl.arange(0, D) + q = tl.load(Q + (hz * L + rows[:, None]) * D + cols[None, :]) + if INT8_QK: + qs = tl.load(QS + hz * L + rows) + nblocks = tl.load(Count + hz * nt + tile) + m = tl.full((64, ), -float("inf"), tl.float32) + den = tl.zeros((64, ), tl.float32) + acc = tl.zeros((64, D), tl.float32) + for block in range(nblocks): + kv = tl.load(Index + (hz * nt + tile) * nt + block) + key_rows = kv * 64 + tl.arange(0, 64) + k = tl.load(K + (hz * L + key_rows[None, :]) * D + cols[:, None]) + if INT8_QK: + ks = tl.load(KS + hz * L + key_rows) + valid = tl.load(VBS + kv) + if valid > 0: + logits = tl.dot(q, k).to(tl.float32) + if INT8_QK: + logits = logits * qs[:, None] * ks[None, :] + logits = logits * (1.4426950408889634 / D**0.5) + logits = tl.where((tl.arange(0, 64) < valid)[None, :], logits, -float("inf")) + block_max = tl.max(logits, 1) + new_m = tl.maximum(m, block_max) + p = tl.exp2(logits - new_m[:, None]) + alpha = tl.exp2(m - new_m) + den = den * alpha + tl.sum(p, 1) + acc = acc * alpha[:, None] + v = tl.load(V + (hz // H) * VB + (hz % H) * VH + key_rows[:, None] * VS_ROW + cols[None, :] * VD) + if FP8_PV: + if P_DYNAMIC: + pscale = tl.maximum(tl.exp2(block_max - new_m) / 448.0, 1e-30) + pv = tl.dot((p / pscale[:, None]).to(tl.float8e4nv), v, out_dtype=tl.float32) + pv = pv * (pscale[:, None] * 448.0) + else: + pv = tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32) + if V_TILE: + scale = tl.load(VS + (hz * nt + kv) * D + cols) + acc += pv * (scale[None, :] / 448.0) + else: + acc += pv + else: + acc += tl.dot(p.to(tl.bfloat16), v, out_dtype=tl.float32) + m = new_m + result = acc / den[:, None] + if FP8_PV and not V_TILE: + vs = tl.load(VS + hz * D + cols) + result = result * (vs[None, :] / 448.0) + result = tl.where(den[:, None] > 0, result, 0.0) + tl.store(Out + (hz * L + rows[:, None]) * D + cols[None, :], result.to(Out.dtype.element_ty)) + + +def sparse_sm89_attention(q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + mask: torch.Tensor, + vbs: torch.Tensor, + *, + int8_qk: bool = True, + fp8_pv: bool = False, + fp8_v_tiles: bool = False, + fp8_dynamic_p: bool = False) -> torch.Tensor: + """Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles.""" + if torch.is_grad_enabled(): + raise ValueError("Sparse INT8/FP8 attention is inference-only") + if not q.is_cuda or torch.cuda.get_device_capability(q.device) != (8, 9): + raise ValueError("Sparse INT8/FP8 attention requires sm89 CUDA") + if q.dtype != torch.bfloat16 or q.shape[-1] != 128 or q.shape != k.shape or q.shape != v.shape: + raise ValueError("Sparse INT8/FP8 attention requires matching BF16 Q/K/V with head dimension 128") + b, h, length, dim = q.shape + if length != vbs.numel() * 64 or mask.shape != (b, h, length // 64, length // 64): + raise ValueError("Sparse INT8/FP8 attention requires a tile-64 mask and validity vector") + from fastvideo_kernel.triton_kernels.index import map_to_index + + # The production INT8-QK/BF16-PV route reads BSHD-backed views directly. + # Quantized Q/K and the output remain contiguous BHSD. Other ablations + # retain their established layout and arithmetic. + if not int8_qk or fp8_pv: + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous() + grid = (triton.cdiv(length, 16), b * h) + qi, ki, vf = q, k, v + qs, ks, vs = q, k, v # unused pointers in BF16 ablations + if int8_qk: + # Tile pads are zero by contract; avoid a full FP32 copy for the reduction. + # Preserve the exact reduction used by the old contiguous adapter; + # its temporary copy dies before Q/K quantization and fine attention. + mean = k.contiguous().sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1) + qi = torch.empty(q.shape, device=q.device, dtype=torch.int8) + ki = torch.empty(k.shape, device=k.device, dtype=torch.int8) + qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32) + ks = torch.empty_like(qs) + _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, h, *q.stride(), CENTER=False, ROWS=16, num_warps=4) + _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, h, *k.stride(), CENTER=True, ROWS=16, num_warps=4) + if fp8_pv: + vf = torch.empty_like(v, dtype=torch.float8_e4m3fn) + if fp8_v_tiles: + vs = torch.empty((b, h, length // 64, dim), device=q.device, dtype=torch.float32) + _quantize_v_tiles[(length // 64, b * h)](v, vf, vs, length, dim, num_warps=8) + else: + vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) + _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) + index, count = map_to_index(mask.contiguous()) + out = torch.empty(q.shape, device=q.device, dtype=q.dtype) + _sparse_int8_fp8[(length // 64, b * h)](qi, + ki, + vf, + qs, + ks, + vs, + index, + count, + vbs, + out, + length, + dim, + h, + *vf.stride(), + INT8_QK=int8_qk, + FP8_PV=fp8_pv, + V_TILE=fp8_v_tiles, + P_DYNAMIC=fp8_dynamic_p) + return out diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index 65401519bd..0ed8e9a647 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -547,6 +547,9 @@ def __init__( self.prefix = prefix self.layer_idx = layer_idx_from_prefix(prefix, default=-1) self.head_size = head_size + self._sm89_kernel = envs.FASTVIDEO_H3_VSA_SM89_KERNEL.get() + if self._sm89_kernel not in {"original", "bf16", "int8"}: + raise ValueError("FASTVIDEO_H3_VSA_SM89_KERNEL must be original, bf16, or int8") # Generic torch.compile must not specialize the shared VSA forward on # the Python ``layer_idx`` value of each of H3's 50 blocks. This # tensor is prepared after weights load and drives only the compiled @@ -778,9 +781,14 @@ def forward( # type: ignore[override] # kernels' granularity. These 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() + sm89_strided = (self._sm89_kernel == "int8" and not torch.is_grad_enabled() and not compiling + and query.dtype == torch.bfloat16 and query.shape[-1] == 128 + and torch.cuda.get_device_capability(query.device) == (8, 9)) + q_bhsd = query.transpose(1, 2) + k_bhsd = key.transpose(1, 2) + v_bhsd = value.transpose(1, 2) + if not sm89_strided: + q_bhsd, k_bhsd, v_bhsd = (t.contiguous() for t in (q_bhsd, k_bhsd, v_bhsd)) sm100a_mask = mask sm100a_variable_block_sizes = attn_metadata.variable_block_sizes @@ -877,19 +885,37 @@ def forward( # type: ignore[override] ) else: if has_sm100a_pair: - q_bhsd = q_bhsd[:, :, :logical_seq_len].contiguous() - k_bhsd = k_bhsd[:, :, :logical_seq_len].contiguous() - v_bhsd = v_bhsd[:, :, :logical_seq_len].contiguous() - out_bhsd, _ = block_sparse_attn_64_bhsd( - q_bhsd, - k_bhsd, - v_bhsd, - mask, - attn_metadata.variable_block_sizes, - ) + q_bhsd = q_bhsd[:, :, :logical_seq_len] + k_bhsd = k_bhsd[:, :, :logical_seq_len] + v_bhsd = v_bhsd[:, :, :logical_seq_len] + if not sm89_strided: + q_bhsd, k_bhsd, v_bhsd = (t.contiguous() for t in (q_bhsd, k_bhsd, v_bhsd)) + if (self._sm89_kernel != "original" and not torch.is_grad_enabled() and not compiling + and q_bhsd.dtype == torch.bfloat16 and q_bhsd.shape[-1] == 128 + and torch.cuda.get_device_capability(q_bhsd.device) == (8, 9)): + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + logger.info_once(f"MiniMax-H3 VSA tile-64 forward: sm89 {self._sm89_kernel} QK / BF16 PV") + out_bhsd = sparse_sm89_attention(q_bhsd, + k_bhsd, + v_bhsd, + mask, + attn_metadata.variable_block_sizes, + int8_qk=self._sm89_kernel == "int8", + fp8_pv=False) + else: + out_bhsd, _ = block_sparse_attn_64_bhsd( + q_bhsd, + k_bhsd, + v_bhsd, + mask, + attn_metadata.variable_block_sizes, + ) if has_sm100a_pair and use_sm100a: out_bhsd = out_bhsd[:, :, :logical_seq_len] out = out_bhsd.transpose(1, 2).contiguous() + # Fine attention is complete. Release its layout copies before the + # gated compression merge creates full-sequence temporaries. + del q_bhsd, k_bhsd, v_bhsd, out_bhsd else: out, _ = block_sparse_attn_256_bshd( logical_query, diff --git a/fastvideo/envs.py b/fastvideo/envs.py index 8962a8e703..77ba871a67 100644 --- a/fastvideo/envs.py +++ b/fastvideo/envs.py @@ -516,6 +516,121 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="eval", doc="Cached reference-feature file for the audio Frechet distance metric.") +# ================== MiniMax-H3 single-GPU switches ================== + +FASTVIDEO_H3_VSA_FP4 = EnvBool(False, + category="attention", + doc="Run MiniMax-H3 VSA attention on the block-sparse SageAttention3 FP4 kernel " + "(sm_120, no-grad, single sequence-parallel rank).") +FASTVIDEO_H3_VSA_TILE_FIRST = EnvBool(False, + category="attention", + doc="Single-rank MiniMax-H3 VSA with one tile gather of the block input " + "instead of separate Q/K/V/gate scatters.") +FASTVIDEO_H3_VSA_SM89_KERNEL = EnvChoice("original", + choices=("original", "bf16", "int8"), + category="attention", + doc="Fine-attention kernel for MiniMax-H3 VSA on sm_89: original, bf16, " + "or int8 (INT8 QK, BF16 PV).") +FASTVIDEO_H3_SIM_SP_FP8 = EnvBool(False, + category="debug", + doc="Simulate the FP8 sequence-parallel exchange of the MiniMax-H3 FP4 VSA " + "path on one rank.") +FASTVIDEO_H3_FFN_CHUNK_TOKENS = EnvInt(0, + category="performance", + doc="Inference-only MiniMax-H3 FFN token chunk size; 0 runs the FFN " + "unchunked.") +FASTVIDEO_H3_FP8_ATTENTION = EnvBool(False, + category="performance", + doc="With NVFP4 layer_profile h3_dit_ffn, run MiniMax-H3 attention " + "projections in FP8.") +FASTVIDEO_H3_FP8_GRANULARITY = EnvChoice("tensor", + choices=("tensor", "channel"), + category="performance", + doc="FP8 scaling granularity for FASTVIDEO_H3_FP8_ATTENTION.") +FASTVIDEO_NVFP4_MM_BACKEND = EnvStr("auto", + category="performance", + doc="FlashInfer mm_fp4 backend for NVFP4 linears, e.g. auto or cutlass.") +FASTVIDEO_NVFP4_ACT_AMAX = EnvPath(None, + category="performance", + doc="JSON of calibrated NVFP4 input amax per linear, keyed b. or " + "full prefix; sets a static activation scale.") +FASTVIDEO_NVFP4_DYNAMIC_ACT = EnvStr("", + category="performance", + doc="NVFP4 linears that derive the activation scale per call: all, or " + "comma-separated layer-name suffixes such as ff.fc_out.") +FASTVIDEO_H3_ADALN_CACHE = EnvBool(False, + category="performance", + doc="Cache MiniMax-H3 AdaLN modulation per timestep instead of keeping the " + "projection weights resident.") +FASTVIDEO_H3_ADALN_TABLE = EnvPath(None, + category="performance", + doc="Precomputed MiniMax-H3 AdaLN modulation table; enables the cache and " + "skips loading the AdaLN projection weights.") +FASTVIDEO_H3_ADALN_DUMP = EnvPath(None, + category="debug", + doc="Write the MiniMax-H3 AdaLN modulation table to this path while sampling.") +FASTVIDEO_H3_SPLICE_TRANSFORMER = EnvPath(None, + category="eval", + doc="Second MiniMax-H3 transformer that runs the late DMD steps " + "(checkpoint step-splice evaluation).") +FASTVIDEO_H3_SPLICE_FROM_STEP = EnvInt(4, + category="eval", + doc="First denoising step run by FASTVIDEO_H3_SPLICE_TRANSFORMER.") +FASTVIDEO_H3_ENCODER_LAYERWISE = EnvBool(False, + category="performance", + doc="Stream MiniMax-H3 text-encoder language layers through exact-size " + "pinned host memory (text-only prompts).") +FASTVIDEO_H3_ENCODER_FUSED_DEQUANT = EnvBool(False, + category="performance", + doc="Expand the serialized NVFP4 MiniMax-H3 text encoder with one " + "fused Triton pass on GPUs without FP4 GEMM.") +FASTVIDEO_H3_VAE_TILE_BATCH = EnvInt(1, + category="performance", + doc="Spatial tiles per MiniMax-H3 video VAE decoder call; 1 decodes per " + "tile.") +FASTVIDEO_H3_VAE_INT8_SHARED_QKV = EnvBool(False, + category="performance", + doc="Share the INT8 activation rotation and quantization across the " + "MiniMax-H3 VAE Q/K/V projections.") +FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW = EnvBool(False, + category="performance", + doc="Use transposed weight views in the MiniMax-H3 VAE INT8 " + "projections.") +FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT = EnvBool(False, + category="performance", + doc="Fused dequantization epilogue for the MiniMax-H3 VAE INT8 " + "projections.") +FASTVIDEO_H3_PINNED_SWAP = EnvBool(True, + category="performance", + doc="Swap offloaded MiniMax-H3 modules through exact-size pinned host " + "arenas.") +FASTVIDEO_H3_PARK_MODULES = EnvStr(None, + category="performance", + doc="Comma-separated MiniMax-H3 denoise modules parked on the host while the " + "text encoder runs, e.g. vae,audio_vae.") +FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS = EnvBool(False, + category="performance", + doc="Layerwise offload also streams large buffers such as packed " + "FP4/FP8 weights.") +FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS = EnvInt(0, + category="performance", + doc="Keep the first N layerwise-offloaded blocks resident on the " + "GPU.") +FASTVIDEO_H3_SP_PROFILE = EnvBool(False, + category="profiling", + doc="CUDA-event spans per stage over one MiniMax-H3 FP4 VSA DiT forward.") +FASTVIDEO_H3_CAPTURE_QKV = EnvPath(None, + category="debug", + doc="Directory for captured real MiniMax-H3 Q/K/V attention inputs.") +FASTVIDEO_CUDA_MEMORY_CAP_GIB = EnvFloat(0.0, + category="debug", + doc="Cap this process's CUDA allocator at this many GiB to emulate a " + "smaller GPU; 0 leaves it uncapped.") +FASTVIDEO_MEMORY_REPORT = EnvBool(False, + category="debug", + doc="Log bytes held per pipeline component by device and dtype after " + "loading.") + # ================== Tests ================== FASTVIDEO_TEST_LTX2_OVERFIT_DATA_DIR = EnvStr("data/cats", diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 05a9aacf54..2d384fd31b 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -1,8 +1,11 @@ from contextlib import contextmanager +from itertools import chain from typing import Any import torch from torch import nn from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager +import fastvideo.envs as envs +from fastvideo.hooks.pinned_memory import PinnedTensorArena from fastvideo.logger import init_logger logger = init_logger(__name__) @@ -14,6 +17,29 @@ def _tensor_placeholder(tensor: torch.Tensor, device: torch.device) -> torch.Ten return torch.empty(shape, device=device, dtype=tensor.dtype) +# Buffers at least this large also stream (e.g. packed NVFP4 weights registered as buffers); +# small ones (scales, caches) stay resident. Opt in with FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS=1. +_BUFFER_OFFLOAD_MIN_BYTES = 1 << 20 + + +def _offload_tensors(module: nn.Module, names: dict[str, torch.Tensor] | None = None): + """``(name, tensor)`` for every parameter and, when enabled, every large buffer. + + ``names`` restricts the walk to the tensors chosen at init: an offloaded buffer is a + zero-element placeholder afterwards and would fail the size test. + """ + if names is not None: + for name, tensor in chain(module.named_parameters(), module.named_buffers()): + if name in names: + yield name, tensor + return + yield from module.named_parameters() + if envs.FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS.get(): + for name, buf in module.named_buffers(): + if buf is not None and buf.numel() * buf.element_size() >= _BUFFER_OFFLOAD_MIN_BYTES: + yield name, buf + + class LayerwiseOffloadState: def __init__( @@ -21,6 +47,7 @@ def __init__( async_copy_stream: torch.cuda.Stream, device: torch.device, next_state: "LayerwiseOffloadState | None" = None, + pin_cpu_memory: bool = True, ) -> None: self.async_copy_stream = async_copy_stream self.next_state = next_state @@ -28,6 +55,8 @@ def __init__( self.cpu_named_parameters: dict[str, torch.Tensor] = {} self.module_ref: nn.Module = None # type: ignore self.device: torch.device = device + self.cpu_arena: PinnedTensorArena | None = None + self.pin_cpu_memory = pin_cpu_memory def _will_offload(self, name: str) -> bool: return True @@ -35,16 +64,33 @@ def _will_offload(self, name: str) -> bool: @torch.compiler.disable def on_init(self, module: nn.Module): self.module_ref = module - for name, param in self.module_ref.named_parameters(): + self.clear_cpu_storage() + if self.pin_cpu_memory: + self.cpu_arena = PinnedTensorArena( + (name, param) for name, param in _offload_tensors(module) if self._will_offload(name)) + for name, param in _offload_tensors(self.module_ref): if self._will_offload(name): - self.cpu_named_parameters[name] = (param.data.detach().to("cpu").pin_memory()) + if self.cpu_arena is not None: + host = self.cpu_arena.empty_like(name, param) + host.copy_(param.data.detach()) + else: + # Retain checkpoint-backed CPU storage so the OS can reclaim + # inactive file pages instead of holding an anonymous pinned copy. + host = param.data.detach().to("cpu") + self.cpu_named_parameters[name] = host param.data = _tensor_placeholder(param.data, self.device) + def clear_cpu_storage(self) -> None: + self.cpu_named_parameters.clear() + if self.cpu_arena is not None: + self.cpu_arena.close() + self.cpu_arena = None + @torch.compiler.disable def wait_and_replace_params(self): torch.cuda.current_stream().wait_stream(self.async_copy_stream) # now gpu_named_parameters are ready - for name, param in self.module_ref.named_parameters(): + for name, param in _offload_tensors(self.module_ref, self.cpu_named_parameters): if not self._will_offload(name): continue if name not in self.gpu_named_parameters: @@ -56,7 +102,7 @@ def wait_and_replace_params(self): def prefetch_params(self): compute_stream = torch.cuda.current_stream() with torch.cuda.stream(self.async_copy_stream): - for name, param in self.module_ref.named_parameters(): + for name, param in _offload_tensors(self.module_ref, self.cpu_named_parameters): if not self._will_offload(name): continue assert name not in self.gpu_named_parameters @@ -66,7 +112,7 @@ def prefetch_params(self): @torch.compiler.disable def release_gpu_params(self): - for name, param in self.module_ref.named_parameters(): + for name, param in _offload_tensors(self.module_ref, self.cpu_named_parameters): if self._will_offload(name): param.data = _tensor_placeholder(param.data, self.device) del self.gpu_named_parameters[name] @@ -83,16 +129,17 @@ def on_attach(self, module: nn.Module): self.state.on_init(module) # pyright: ignore def on_detach(self, module: nn.Module): - named_parameters = dict(module.named_parameters()) + self.state.async_copy_stream.synchronize() + named_parameters = dict(_offload_tensors(module, self.state.cpu_named_parameters)) for name, cpu_tensor in self.state.cpu_named_parameters.items(): - if name not in self.state.gpu_named_parameters: - if name in named_parameters: - named_parameters[name].data = cpu_tensor.to(device=self.state.device) - else: - logger.warning( - "Parameter {} not found in module during detachment.", - name, - ) + if name in named_parameters: + gpu_tensor = self.state.gpu_named_parameters.get(name) + named_parameters[name].data = gpu_tensor if gpu_tensor is not None else cpu_tensor.to(self.state.device) + else: + logger.warning("Parameter %s not found in module during detachment.", name) + self.state.gpu_named_parameters.clear() + self.state.clear_cpu_storage() + self.state.next_state = None @classmethod def name(cls) -> str: @@ -129,12 +176,21 @@ def mutate_params_scope(self): yield finally: # instead of releasing, we should overwrite the original params since they have been modified - self.state.cpu_named_parameters.clear() self.state.gpu_named_parameters.clear() self.state.on_init(self.state.module_ref) # pyright: ignore -def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): +def enable_layerwise_offload(model: nn.Module, + is_replace: bool = False, + *, + resident_blocks: int | None = None, + cyclic: bool = True, + pin_cpu_memory: bool = True): + """Stream blocks, optionally retaining their existing pageable CPU storage. + + Disabling pinning avoids a private copy of file-backed checkpoint tensors. + Transfers can be slower, but inactive checkpoint pages remain reclaimable. + """ if torch.cuda.is_available(): device = torch.device("cuda", torch.cuda.current_device()) else: @@ -142,10 +198,24 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): return state_list = [] async_stream = torch.cuda.Stream() + # The first N entries skip offloading and stay wherever the model is placed (normally the + # GPU), so a GPU with spare memory streams only the remainder over PCIe. + if resident_blocks is not None: + resident = max(0, resident_blocks) + else: + try: + resident = max(0, envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.get()) + except ValueError as error: + logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS: %s", error) + resident = 0 for name, submodule in model.named_children(): if isinstance(submodule, nn.ModuleList): for idx, module_entry in enumerate(submodule): - state = LayerwiseOffloadState(async_copy_stream=async_stream, device=device) + if idx < resident: + continue + state = LayerwiseOffloadState(async_copy_stream=async_stream, + device=device, + pin_cpu_memory=pin_cpu_memory) state_list.append(state) hook_mgr = ModuleHookManager.get_from_or_default(module_entry) hook = LayerwiseOffloadHook(state) @@ -159,8 +229,13 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): hook_mgr.append_forward_hook(hook) break if len(state_list) == 0: + if resident > 0: + logger.info("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%d keeps every block resident; nothing to offload", + resident) + return raise ValueError("No nn.ModuleList found in the model for layerwise offloading.") - # circular linking of states + # Repeated DiT steps prefetch the first block after the last. A once-per-request + # encoder can skip that unused copy and release every layer after its forward. for i in range(len(state_list)): - state_list[i].next_state = state_list[(i + 1) % len(state_list)] + state_list[i].next_state = state_list[(i + 1) % len(state_list)] if cyclic or i + 1 < len(state_list) else None diff --git a/fastvideo/hooks/pinned_memory.py b/fastvideo/hooks/pinned_memory.py new file mode 100644 index 0000000000..81e7ff8af3 --- /dev/null +++ b/fastvideo/hooks/pinned_memory.py @@ -0,0 +1,86 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exact-size pinned host storage for inference offload.""" + +import weakref +from collections.abc import Iterable + +import torch + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) +_ALIGNMENT = 256 +_PAGE_ALIGNMENT = 4096 + + +def _unregister(buffer: torch.Tensor, device: int) -> None: + # The allocation must outlive any outstanding nonblocking H2D copies. + try: + with torch.cuda.device(device): + torch.cuda.synchronize() + error = torch.cuda.cudart().cudaHostUnregister(buffer.data_ptr()) + if error != 0: + logger.warning("cudaHostUnregister failed: %s", error) + except Exception as exc: + # CUDA may already be unavailable during interpreter shutdown. + logger.warning("Could not unregister pinned host arena: %s", exc) + + +class PinnedTensorArena: + """Pack tensors into one CUDA-registered CPU allocation, aligned to 256 bytes. + + PyTorch's pinned allocator rounds large allocations to powers of two. Registering + ordinary host storage avoids that overhead. Typed views retain this owner, so + registration survives even if the module or offload state is dropped first. + If registration is unavailable, allocate conventional pinned tensors instead. + Call ``close`` only after all views and pending copies have been released. + """ + + def __init__(self, tensors: Iterable[tuple[str, torch.Tensor]]) -> None: + self.offsets: dict[str, tuple[int, int]] = {} + size = 0 + for name, tensor in tensors: + size = (size + _ALIGNMENT - 1) // _ALIGNMENT * _ALIGNMENT + length = tensor.numel() * tensor.element_size() + self.offsets[name] = (size, length) + size += length + self.nbytes = size + self.buffer: torch.Tensor | None = None + self._finalizer: weakref.finalize | None = None + if not size: + return + # Register dedicated pages: small malloc allocations can otherwise share a + # registered page with another arena. The extra space is bounded by 8 KiB. + span = (size + _PAGE_ALIGNMENT - 1) // _PAGE_ALIGNMENT * _PAGE_ALIGNMENT + allocation = torch.empty(span + _PAGE_ALIGNMENT - 1, dtype=torch.uint8, device="cpu") + start = (-allocation.data_ptr()) % _PAGE_ALIGNMENT + # Give the aligned region its own storage base. Tensor.is_pinned() queries + # the storage pointer, which would precede the registered region for a + # plain narrow() view. The memoryview retains the original allocation. + buffer = torch.frombuffer(memoryview(allocation.numpy())[start:start + span], dtype=torch.uint8) + device = torch.cuda.current_device() + try: + error = torch.cuda.cudart().cudaHostRegister(buffer.data_ptr(), span, 0) + if error != 0: + raise RuntimeError(f"cudaHostRegister returned {error}") + except Exception as exc: + logger.warning("Exact-size host registration failed; using the pinned allocator: %s", exc) + return + self.buffer = buffer + self._finalizer = weakref.finalize(self, _unregister, buffer, device) + + def empty_like(self, name: str, tensor: torch.Tensor) -> torch.Tensor: + """Return a contiguous host view with the source's dtype and shape.""" + if self.buffer is None: + return torch.empty(tensor.shape, dtype=tensor.dtype, device="cpu", pin_memory=True) + offset, length = self.offsets[name] + host = self.buffer.narrow(0, offset, length).view(tensor.dtype).reshape(tensor.shape) + host._pinned_arena = self + return host + + def close(self) -> None: + """Unregister before releasing storage; safe to call more than once.""" + if self._finalizer is not None: + self._finalizer() + self._finalizer = None + self.buffer = None diff --git a/fastvideo/layers/quantization/fp8_config.py b/fastvideo/layers/quantization/fp8_config.py index e73c372d65..5b7e3f3cbf 100644 --- a/fastvideo/layers/quantization/fp8_config.py +++ b/fastvideo/layers/quantization/fp8_config.py @@ -3,6 +3,8 @@ Matches linear layers by suffix (``to_q/k/v/to_out``, ``ffn.fc_in/fc_out``). Supports per-tensor (default, fast) and per-channel (higher accuracy) granularity. +On sm89 the per-channel path runs the per-tensor GEMM plus a Triton scale epilogue +(``fp8_kernels``), since torch's rowwise-scaled kernel there is slower than bf16. Falls back to bf16 dequant on GPUs older than sm89. """ from __future__ import annotations @@ -32,6 +34,9 @@ _FP8_SUFFIXES = ( "ffn.fc_in", "ffn.fc_out", + # MiniMax-H3 blocks name their MLP ``ff``. + "ff.fc_in", + "ff.fc_out", "to_q", "to_k", "to_v", @@ -74,8 +79,18 @@ def _quantize_tensorwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor return x_fp8, x_scale.view(1) +def _rowwise_scaled_mm_is_slow() -> bool: + if getattr(torch.version, "hip", None) or not torch.cuda.is_available(): + return False + from fastvideo.layers.quantization.fp8_kernels import rowwise_scaled_mm_is_slow + return rowwise_scaled_mm_is_slow() + + def _quantize_rowwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``.""" + if x_2d.is_cuda and _rowwise_scaled_mm_is_slow(): + from fastvideo.layers.quantization.fp8_kernels import quantize_rowwise_fp8 + return quantize_rowwise_fp8(x_2d) x_absmax = x_2d.abs().amax(dim=-1, keepdim=True).float() x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE) x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE) @@ -156,6 +171,12 @@ def apply( w_fp8 = layer._fp8_weight w_scale = layer._fp8_weight_scale + if self.granularity == "channel" and _rowwise_scaled_mm_is_slow(): + from fastvideo.layers.quantization.fp8_kernels import scaled_mm_token_channel + out = scaled_mm_token_channel(x_fp8, x_scale, w_fp8.t(), w_scale) + if bias is not None: + out = out + bias + return out.view(*original_shape[:-1], out_dim) scale_b = w_scale.view(1, -1) if self.granularity == "channel" else w_scale out = torch._scaled_mm( diff --git a/fastvideo/layers/quantization/fp8_kernels.py b/fastvideo/layers/quantization/fp8_kernels.py new file mode 100644 index 0000000000..a589262c3a --- /dev/null +++ b/fastvideo/layers/quantization/fp8_kernels.py @@ -0,0 +1,84 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Triton helpers for per-token x per-channel FP8 linears on GPUs whose rowwise ``_scaled_mm`` is slow. + +On sm89 (RTX 4090 / L40S / RTX 6000 Ada) ``torch._scaled_mm`` with rowwise scales runs at +~70 TFLOPS, below bf16 (~160), while the per-tensor kernel reaches 220-305 TFLOPS. The +per-token x per-channel result is recovered exactly by running the per-tensor kernel with unit +scales and applying ``out[i, j] *= sx[i] * sw[j]`` in one pass over the output, which costs +5-10% of the GEMM instead of 2-4x. +""" +from __future__ import annotations + +import functools + +import torch +import triton +import triton.language as tl + +FP8_MAX = 448.0 +FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0) + + +@functools.cache +def rowwise_scaled_mm_is_slow() -> bool: + """Ada (sm89) has no fast rowwise-scaled FP8 GEMM in torch; Hopper and Blackwell do.""" + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (8, 9) + + +@triton.jit +def _quantize_rowwise_kernel(x_ptr, q_ptr, s_ptr, K, stride_x, stride_q, BLOCK_K: tl.constexpr): + row = tl.program_id(0).to(tl.int64) + x_row = x_ptr + row * stride_x + amax = tl.zeros((BLOCK_K, ), dtype=tl.float32) + for k in range(0, K, BLOCK_K): + cols = k + tl.arange(0, BLOCK_K) + amax = tl.maximum(amax, tl.abs(tl.load(x_row + cols, mask=cols < K, other=0.0).to(tl.float32))) + scale = tl.maximum(tl.max(amax, axis=0) / 448.0, 1.0 / (448.0 * 512.0)) + tl.store(s_ptr + row, scale) + inv = 1.0 / scale + for k in range(0, K, BLOCK_K): + cols = k + tl.arange(0, BLOCK_K) + v = tl.load(x_row + cols, mask=cols < K, other=0.0).to(tl.float32) * inv + v = tl.minimum(tl.maximum(v, -448.0), 448.0) + tl.store(q_ptr + row * stride_q + cols, v.to(tl.float8e4nv), mask=cols < K) + + +def quantize_rowwise_fp8(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Per-token FP8 quantization in one launch. Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``.""" + x_2d = x_2d.contiguous() + M, K = x_2d.shape + q = torch.empty((M, K), device=x_2d.device, dtype=torch.float8_e4m3fn) + s = torch.empty((M, 1), device=x_2d.device, dtype=torch.float32) + if M: + _quantize_rowwise_kernel[(M, )](x_2d, q, s, K, x_2d.stride(0), q.stride(0), BLOCK_K=1024, num_warps=4) + return q, s + + +@triton.jit +def _scale_rows_cols_kernel(o_ptr, sx_ptr, sw_ptr, M, N, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr): + rows = tl.program_id(0).to(tl.int64) * BLOCK_M + tl.arange(0, BLOCK_M) + cols = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N) + mask = (rows[:, None] < M) & (cols[None, :] < N) + # int64 offsets: a 78k-token fc_in output has 2.2e9 elements, past int32. + ptrs = o_ptr + rows[:, None] * N + cols[None, :] + v = tl.load(ptrs, mask=mask, other=0.0).to(tl.float32) + sx = tl.load(sx_ptr + rows, mask=rows < M, other=0.0) + sw = tl.load(sw_ptr + cols, mask=cols < N, other=0.0) + tl.store(ptrs, (v * sx[:, None] * sw[None, :]).to(o_ptr.dtype.element_ty), mask=mask) + + +def scaled_mm_token_channel(x_fp8: torch.Tensor, x_scale: torch.Tensor, w_fp8_t: torch.Tensor, + w_scale: torch.Tensor) -> torch.Tensor: + """``(x_fp8 * x_scale) @ (w_fp8_t * w_scale)`` in bf16 via the fast per-tensor GEMM plus a scale epilogue. + + The unit-scale GEMM output is at most 448^2 * K, far inside bf16 range, and its relative + precision is that of any bf16 output, so the epilogue loses nothing against rowwise scaling. + """ + one = torch.ones((), device=x_fp8.device, dtype=torch.float32) + out = torch._scaled_mm(x_fp8, w_fp8_t, scale_a=one, scale_b=one, out_dtype=torch.bfloat16) + if isinstance(out, tuple): + out = out[0] + M, N = out.shape + grid = (triton.cdiv(M, 64), triton.cdiv(N, 128)) + _scale_rows_cols_kernel[grid](out, x_scale.reshape(-1), w_scale.reshape(-1), M, N, BLOCK_M=64, BLOCK_N=128) + return out diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 284622d595..3e683f749c 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -8,8 +8,9 @@ variants that may land later (e.g. AMD's MX-FP4 or vendor-neutral e3m0). -The registered config targets the curated LTX-2 deployment set and the -main MiniMax-H3 transformer-block FFN linears. +The registered config targets the curated LTX-2 deployment set, the +main MiniMax-H3 transformer-block FFN linears, and the packed MiniMax-H3 +DiT export (``layer_profile="h3_dit"``) covering attention plus FFN. `flashinfer` is imported lazily inside the call paths that need it. This keeps ``import fastvideo`` cheap on hosts where flashinfer is @@ -19,6 +20,7 @@ from __future__ import annotations import logging +import os import re from typing import Any @@ -26,6 +28,8 @@ import torch.nn.functional as F from torch.nn.parameter import Parameter +import fastvideo.envs as envs + from fastvideo.layers.quantization.base_config import ( QuantizationConfig, QuantizeMethodBase, @@ -76,6 +80,21 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: for suffix in _LTX2_NVFP4_BLOCK_LINEAR_SUFFIXES) | frozenset( ("ltx2.adaln_single.linear", )) _MINIMAX_H3_NVFP4_FF_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.ff\.(?:fc_in|fc_out)$") +_MINIMAX_H3_NVFP4_DIT_PREFIX = re.compile( + r"(?:^|\.)transformer_blocks\.\d+\.(?:attn\.to_(?:q|k|v|out)|ff\.(?:fc_in|fc_out))$") +_H3_BLOCK_ATTN_PROJ = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out)$") +_MINIMAX_H3_NVFP4_VSA_GATE_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.attn\.to_gate_compress$") +H3_NVFP4_DIT_EXPORT_FILENAME = "nvfp4_weights.safetensors" +H3_NVFP4_DIT_KEY_SEP = "::" +H3_NVFP4_DIT_BUFFER_NAMES = ( + "_nvfp4_weight", + "_nvfp4_weight_scale", + "_nvfp4_alpha", + "_weight_global_sf", +) +# Optional per-layer static activation global scale (448 * 6 / calibrated input amax). +# Exports without it quantize activations with the unit global scale. +H3_NVFP4_DIT_INPUT_SF_NAME = "_nvfp4_input_global_sf" def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: @@ -84,10 +103,43 @@ def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: def is_minimax_h3_nvfp4_linear_prefix(prefix: str) -> bool: - """Return whether *prefix* is a main MiniMax-H3 transformer-block FFN linear.""" return _MINIMAX_H3_NVFP4_FF_PREFIX.search(prefix) is not None +def is_minimax_h3_nvfp4_dit_linear_prefix(prefix: str) -> bool: + """Return whether *prefix* is a MiniMax-H3 DiT attention or FFN linear. + + This is the packed NVFP4H3 export set: ``attn.to_{q,k,v,out}`` and + ``ff.{fc_in,fc_out}`` in each main transformer block. Token-refiner, + AdaLN, and embedding linears stay dense. + """ + return _MINIMAX_H3_NVFP4_DIT_PREFIX.search(prefix) is not None + + +def is_minimax_h3_nvfp4_dit_export_path(path: str) -> bool: + return os.path.basename(path) == H3_NVFP4_DIT_EXPORT_FILENAME + + +def find_minimax_h3_nvfp4_dit_export(weight_paths: list[str]) -> str | None: + seen: list[str] = [] + for path in weight_paths: + if is_minimax_h3_nvfp4_dit_export_path(path) and os.path.isfile(path): + return path + directory = path if os.path.isdir(path) else os.path.dirname(path) + if directory and directory not in seen: + seen.append(directory) + for directory in seen: + candidate = os.path.join(directory, H3_NVFP4_DIT_EXPORT_FILENAME) + if os.path.isfile(candidate): + return candidate + return None + + +def dense_transformer_safetensors(weight_paths: list[str]) -> list[str]: + """Drop the packed NVFP4 DiT export so it is not loaded as bf16 weights.""" + return [path for path in weight_paths if not is_minimax_h3_nvfp4_dit_export_path(path)] + + def _is_ltx2_refine_only_prefix(prefix: str) -> bool: return any(prefix.endswith(suffix) for suffix in _LTX2_REFINE_ONLY_SUFFIXES) @@ -290,6 +342,16 @@ def _mm_fp4( ) +def _mm_fp4_backend() -> str: + """FlashInfer ``mm_fp4`` backend (``FASTVIDEO_NVFP4_MM_BACKEND``, default ``auto``). + + On sm_120 ``auto`` picks a kernel about 2x slower than ``cutlass`` or + ``cudnn`` once activations reach tens of thousands of rows (measured at + 73k rows on an RTX PRO 6000); short sequences are unaffected. + """ + return envs.FASTVIDEO_NVFP4_MM_BACKEND.get() + + def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: """Coerce an activation to a dtype the FP4 linear accepts. @@ -306,8 +368,24 @@ def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: return x +_AMAX_TABLES: dict[str, dict[str, float]] = {} + + +def _load_amax_table(path: str) -> dict[str, float]: + if path not in _AMAX_TABLES: + import json + with open(path) as f: + raw = json.load(f) + _AMAX_TABLES[path] = {k: float(v["all"] if isinstance(v, dict) else v) for k, v in raw.items()} + return _AMAX_TABLES[path] + + class NVFP4QuantizeMethod(QuantizeMethodBase): + # Lazily resolved by _static_activation_global_sf; class defaults also cover object.__new__ test doubles. + _static_sf_checked: bool = False + _static_sf: torch.Tensor | None = None + def __init__(self, layer_prefix: str = ""): super().__init__() self.weight_fp4 = None @@ -333,6 +411,42 @@ def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, layer.register_parameter("weight", weight) set_weight_attrs(weight, extra_weight_attrs) + def _static_activation_global_sf(self) -> torch.Tensor | None: + """FASTVIDEO_NVFP4_ACT_AMAX: JSON of calibrated input amax per layer ("b." or full prefix).""" + if getattr(self, "_static_sf_checked", False): + return self._static_sf + self._static_sf_checked, self._static_sf = True, None + path = envs.FASTVIDEO_NVFP4_ACT_AMAX.get() + if path: + table = _load_amax_table(path) + prefix = self.layer_prefix or "" + match = re.search(r"transformer_blocks\.(\d+)\.(.+)$", prefix) + keys = [prefix] + ([f"b{match.group(1)}.{match.group(2)}"] if match else []) + amax = next((table[k] for k in keys if k in table), None) + if amax is not None: + self._static_sf = torch.tensor((448.0 * 6.0) / max(amax, 1e-12), dtype=torch.float32, device="cuda") + return self._static_sf + + def _dynamic_activation_scale(self) -> bool: + """FASTVIDEO_NVFP4_DYNAMIC_ACT: "all", or comma-separated layer-name suffixes (e.g. "ff.fc_out").""" + cached = getattr(self, "_dynamic_act_cached", None) + if cached is None: + selected = envs.FASTVIDEO_NVFP4_DYNAMIC_ACT.get() + suffixes = [part.strip() for part in selected.split(",") if part.strip()] + prefix = self.layer_prefix or "" + cached = "all" in suffixes or any(prefix.endswith(suffix) for suffix in suffixes) + self._dynamic_act_cached = cached + return cached + + def uses_unit_activation_scale(self, layer: torch.nn.Module) -> bool: + """Whether ``apply`` quantizes this layer's input with the unit global scale. + + False when a calibrated scale (env table or the export's ``_nvfp4_input_global_sf``) or a dynamic + per-call scale applies; such inputs cannot share one pre-quantized copy across layers. + """ + return (self._static_activation_global_sf() is None and getattr(layer, H3_NVFP4_DIT_INPUT_SF_NAME, None) is None + and not self._dynamic_activation_scale()) + def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: SfLayout, _, _ = _require_flashinfer() x = _coerce_fp4_input_dtype(x) @@ -392,7 +506,17 @@ def apply( else: x = _coerce_fp4_input_dtype(x) x = x.view(-1, x.shape[-1]) - x_global_sf = self.x_global_sf + static_sf = self._static_activation_global_sf() + if static_sf is None: + static_sf = getattr(layer, H3_NVFP4_DIT_INPUT_SF_NAME, None) + if static_sf is not None: + x_global_sf = static_sf + elif self._dynamic_activation_scale(): + # A unit global scale caps FP8 block scales at |x| = 6 * 448; inputs such as H3's ff.fc_out + # (post-SwiGLU) exceed that, so derive the global scale from this call's amax. + x_global_sf = (448.0 * 6.0) / x.abs().amax().float().clamp(min=1e-12) + else: + x_global_sf = self.x_global_sf x_fp4, x_scale = _nvfp4_quantize( x, x_global_sf, @@ -417,7 +541,7 @@ def apply( alpha, torch.bfloat16, None, - backend='auto', + backend=_mm_fp4_backend(), ) if bias is not None: @@ -431,14 +555,22 @@ class NVFP4Config(QuantizationConfig): NVFP4 is NVIDIA's block-scaled FP4 (e2m1 mantissa, fp32 alpha, ``layout_128x4`` scale layout, group size 16). LTX-2 uses its curated - attention and FFN deployment set. MiniMax-H3 uses only ``fc_in`` and - ``fc_out`` in each main transformer-block FFN. + attention and FFN deployment set. MiniMax-H3's default profile uses only + ``fc_in`` and ``fc_out`` in each main transformer-block FFN. + ``layer_profile="h3_dit"`` expands that to the packed NVFP4H3 DiT set + (attention ``to_{q,k,v,out}`` plus those FFN linears). + ``layer_profile="h3_dit_ffn"`` loads a packed export of the FFN linears + only, keeping attention projections dense (e.g. calibrated FFN-only + checkpoints such as FastH3 V2 NVFP4). + ``layer_profile="h3_dit_vsa"`` is ``h3_dit`` plus each block's VSA + compression gate ``attn.to_gate_compress`` (VSA-distilled students). """ def __init__(self, layer_profile: str = "refine", retain_original_weights: bool | None = None): super().__init__() - # ``base``: stage-1 set (no attn2.to_out, no cross-modal AV - # projections). ``refine``: full stage-2 set. + if layer_profile not in ("base", "refine", "h3_dit", "h3_dit_ffn", "h3_dit_vsa"): + raise ValueError("NVFP4Config.layer_profile must be one of 'base', 'refine', 'h3_dit', " + f"'h3_dit_ffn', or 'h3_dit_vsa', got {layer_profile!r}") self.layer_profile = layer_profile # Original bf16 ``layer.weight`` retention after FP4 conversion. # Default (None/False): purge the purgeable originals -- every @@ -473,13 +605,26 @@ def from_config(cls, config: dict[str, Any]) -> NVFP4Config: def get_quant_method(self, layer: torch.nn.Module, prefix: str): from fastvideo.layers.linear import LinearBase - # LTX-2 switches its active subset by stage at runtime. MiniMax-H3 - # uses the fixed main-transformer FFN set selected by its prefix. - if isinstance(layer, LinearBase) and (is_ltx2_nvfp4_linear_prefix(prefix) - or is_minimax_h3_nvfp4_linear_prefix(prefix)): + if not isinstance(layer, LinearBase): + return None + if self.layer_profile == "h3_dit": + tagged = is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + elif self.layer_profile == "h3_dit_vsa": + tagged = (is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + or _MINIMAX_H3_NVFP4_VSA_GATE_PREFIX.search(prefix) is not None) + elif self.layer_profile == "h3_dit_ffn": + tagged = is_minimax_h3_nvfp4_linear_prefix(prefix) + else: + tagged = is_ltx2_nvfp4_linear_prefix(prefix) or is_minimax_h3_nvfp4_linear_prefix(prefix) + if tagged: method = NVFP4QuantizeMethod(layer_prefix=prefix) method._retain_original_weights = self.retain_original_weights return method + if (self.layer_profile == "h3_dit_ffn" and envs.FASTVIDEO_H3_FP8_ATTENTION.get() + and _H3_BLOCK_ATTN_PROJ.search(prefix) is not None): + # Mixed precision: NVFP4 MLPs, FP8 (per-tensor weight, dynamic per-tensor activation) attention. + from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod + return FP8QuantizeMethod(granularity=envs.FASTVIDEO_H3_FP8_GRANULARITY.get()) return None @@ -534,9 +679,8 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: if retain: retained += 1 elif isinstance(weight, DTensor): - # ponytail: purging FSDP-sharded originals needs per-shard - # resharding bookkeeping; skip until a sharded deploy needs it. - retained += 1 + raise RuntimeError("NVFP4 cannot purge FSDP-sharded bf16 weights. Use a packed NVFP4 " + "export, or convert without FSDP sharding.") else: purged_bytes += weight.numel() * weight.element_size() purged += 1 @@ -553,10 +697,95 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: ) +def nvfp4_linear_weight_param_names(model: torch.nn.Module) -> set[str]: + """State-dict names of ``weight`` on layers tagged with ``NVFP4QuantizeMethod``.""" + names: set[str] = set() + for module_name, module in model.named_modules(): + if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod): + names.add(f"{module_name}.weight" if module_name else "weight") + return names + + +def _module_by_nvfp4_export_prefix(modules: dict[str, torch.nn.Module], prefix: str) -> torch.nn.Module | None: + module = modules.get(prefix) + if module is not None: + return module + if prefix.startswith("minimax_h3."): + return modules.get(prefix[len("minimax_h3."):]) + return modules.get(f"minimax_h3.{prefix}") + + +def load_minimax_h3_nvfp4_dit_export( + model: torch.nn.Module, + path: str, + device: torch.device | str, +) -> int: + """Load a packed NVFP4H3 DiT export onto already-tagged NVFP4 linears. + + Keys are ``::`` with the four buffers + ``convert_model_to_nvfp4`` registers, plus an optional calibrated + ``_nvfp4_input_global_sf``. Every export prefix must match an + NVFP4 linear, and every NVFP4 linear must appear in the export. + """ + from safetensors import safe_open + + modules = dict(model.named_modules()) + tagged = { + name + for name, module in modules.items() if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod) + } + groups: dict[str, dict[str, str]] = {} + with safe_open(path, framework="pt", device="cpu") as reader: + for key in reader.keys(): # noqa: SIM118 + if H3_NVFP4_DIT_KEY_SEP not in key: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export key {key!r} is missing {H3_NVFP4_DIT_KEY_SEP!r}") + prefix, buffer_name = key.split(H3_NVFP4_DIT_KEY_SEP, 1) + groups.setdefault(prefix, {})[buffer_name] = key + + loaded_names: set[str] = set() + for prefix, buffers in groups.items(): + missing_buffers = [name for name in H3_NVFP4_DIT_BUFFER_NAMES if name not in buffers] + if missing_buffers: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} is missing {missing_buffers}") + extra_buffers = sorted(set(buffers) - set(H3_NVFP4_DIT_BUFFER_NAMES) - {H3_NVFP4_DIT_INPUT_SF_NAME}) + if extra_buffers: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} has unknown buffers {extra_buffers}") + module = _module_by_nvfp4_export_prefix(modules, prefix) + if module is None: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} is not in the model") + if not isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod): + raise RuntimeError("MiniMax-H3 NVFP4 DiT export layer " + f"{prefix!r} is not an NVFP4 linear; set NVFP4Config(layer_profile='h3_dit')") + for buffer_name in H3_NVFP4_DIT_BUFFER_NAMES + (H3_NVFP4_DIT_INPUT_SF_NAME, ): + if buffer_name not in buffers: + continue + tensor = reader.get_tensor(buffers[buffer_name]).to(device=device) + module.register_buffer(buffer_name, tensor, persistent=False) + module.register_parameter("weight", None) + loaded_names.add(next(name for name, candidate in modules.items() if candidate is module)) + + missing_layers = tagged - loaded_names + extra_layers = loaded_names - tagged + if missing_layers or extra_layers: + raise RuntimeError("MiniMax-H3 NVFP4 DiT export does not cover the tagged linear set; " + f"missing={sorted(missing_layers)[:8]} extra={sorted(extra_layers)[:8]}") + logger.info("Loaded MiniMax-H3 NVFP4 DiT export: %d linears from %s", len(loaded_names), path) + return len(loaded_names) + + __all__ = [ + "H3_NVFP4_DIT_BUFFER_NAMES", + "H3_NVFP4_DIT_EXPORT_FILENAME", + "H3_NVFP4_DIT_INPUT_SF_NAME", "NVFP4Config", "NVFP4QuantizeMethod", "convert_model_to_nvfp4", + "dense_transformer_safetensors", + "find_minimax_h3_nvfp4_dit_export", "is_ltx2_nvfp4_linear_prefix", + "is_minimax_h3_nvfp4_dit_export_path", + "is_minimax_h3_nvfp4_dit_linear_prefix", "is_minimax_h3_nvfp4_linear_prefix", + "load_minimax_h3_nvfp4_dit_export", + "nvfp4_linear_weight_param_names", ] diff --git a/fastvideo/layers/quantization/nvfp4_dequant.py b/fastvideo/layers/quantization/nvfp4_dequant.py new file mode 100644 index 0000000000..5bc3c2b342 --- /dev/null +++ b/fastvideo/layers/quantization/nvfp4_dequant.py @@ -0,0 +1,54 @@ +# SPDX-License-Identifier: Apache-2.0 +"""One-pass serialized NVFP4 weight expansion for BF16 consumer-GPU compute.""" +import torch +import triton +import triton.language as tl + + +@triton.jit +def _dequantize_nvfp4(P, S, OUT, N: tl.constexpr, K: tl.constexpr, INVERSE_SCALE, BLOCK: tl.constexpr): + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + valid = offsets < N * K + row, col = offsets // K, offsets % K + packed = tl.load(P + row * (K // 2) + col // 2, valid, other=0).to(tl.uint32) + code = (packed >> ((col % 2) * 4)) & 15 + magnitude = (code & 7).to(tl.float32) + value = tl.where(magnitude < 4, magnitude * 0.5, tl.where(magnitude < 6, magnitude - 2, magnitude * 2 - 8)) + value = value * tl.where((code & 8) != 0, -1.0, 1.0) + group = col // 16 + # FlashInfer layout_128x4: [row_tile, col_tile, row%32, row//32%4, col%4]. + scale_index = ((((row // 128) * (K // 64) + group // 4) * 32 + row % 32) * 4 + (row // 32) % 4) * 4 + group % 4 + scale = tl.load(S + scale_index, valid, other=0.0).to(tl.float32) + output = (value * scale) * INVERSE_SCALE + tl.store(OUT + offsets, output, valid) + + +def dequantize_nvfp4_cuda(packed: torch.Tensor, + scales: torch.Tensor, + global_scale: float, + dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + """Expand E2M1 nibbles and swizzled E4M3 scales without full FP32 intermediates.""" + if not packed.is_cuda or scales.device != packed.device: + raise ValueError("NVFP4 fused dequantization requires tensors on the same CUDA device") + if packed.ndim != 2 or packed.dtype != torch.uint8 or scales.dtype != torch.uint8: + raise ValueError("NVFP4 fused dequantization requires packed uint8 weights and scales") + if not packed.is_contiguous() or not scales.is_contiguous(): + raise ValueError("NVFP4 fused dequantization requires contiguous tensors") + rows, cols = packed.shape[0], packed.shape[1] * 2 + if rows % 128 or cols % 64 or scales.numel() != rows * cols // 16: + raise ValueError("NVFP4 fused dequantization requires exact 128x4 scale geometry") + if dtype not in (torch.bfloat16, torch.float16, torch.float32): + raise ValueError("NVFP4 fused dequantization requires a floating output dtype") + output = torch.empty((rows, cols), dtype=dtype, device=packed.device) + _dequantize_nvfp4[(triton.cdiv(rows * cols, 1024), )]( + packed, + scales.view(torch.float8_e4m3fn), + output, + rows, + cols, + # Match Torch's CPU-scalar division: form the + # reciprocal in double, then cast to FP32. + 1.0 / global_scale, + BLOCK=1024, + num_warps=4) + return output diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index e25650466d..dd9a645a1e 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -4,6 +4,7 @@ from __future__ import annotations import contextlib +import json import math from collections.abc import Iterator, Mapping from typing import Any @@ -14,6 +15,8 @@ from fastvideo import envs from fastvideo.attention import DistributedAttention +from fastvideo.attention.backends.abstract import layer_idx_from_prefix +from fastvideo.attention.backends.video_sparse_attn_h3 import MiniMaxH3VSAMetadata from fastvideo.attention.layer import DistributedAttention_VSA from fastvideo.attention.selector import get_attn_backend from fastvideo.configs.models.dits.minimax_h3 import MiniMaxH3Config @@ -21,7 +24,9 @@ sequence_model_parallel_all_gather_with_unpad, sequence_model_parallel_shard, ) -from fastvideo.distributed.parallel_state import get_sp_world_size, model_parallel_is_initialized +from fastvideo.distributed.parallel_state import (get_sp_group, get_sp_world_size, + model_parallel_is_initialized) +from fastvideo.forward_context import get_forward_context from fastvideo.layers.linear import ReplicatedLinear from fastvideo.layers.quantization.mxfp8_config import MXFP8QuantizeMethod from fastvideo.layers.mlp import MLP @@ -30,6 +35,8 @@ 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_vsa_fp4 import (STAGES, vsa_fp4_attention, vsa_fp4_attention_sp, + vsa_fp4_requested, vsa_tile_first_attention) from fastvideo.models.dits.minimax_h3_fusions import ( HAVE_TRITON, fused_qknorm_rope, @@ -74,6 +81,13 @@ def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool: return HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() +@torch.compile(dynamic=True, fullgraph=True) +def _gated_residual(hidden_states: torch.Tensor, gate_table: torch.Tensor, indices: torch.Tensor, + branch: torch.Tensor) -> torch.Tensor: + """``hidden + gate[indices] * branch`` in one pass (eager materializes the gathered gate).""" + return hidden_states + gate_table.index_select(0, indices) * branch + + class MiniMaxH3RotaryPosEmbed(nn.Module): """Three-axis rotary frequencies over packed `(t, h, w)` coordinates.""" @@ -125,8 +139,22 @@ def __init__( self.fuse_swiglu = fuse_swiglu self.use_mxfp8 = isinstance(self.fc_in.quant_method, MXFP8QuantizeMethod) and isinstance( self.fc_out.quant_method, MXFP8QuantizeMethod) + # Inference-only token chunking: the 2 * ffn_dim intermediate is ~5.3x the block input + # (4.5 GiB at 78k tokens), so chunks bound the activation peak on 24-32 GB GPUs. + self.chunk_tokens = envs.FASTVIDEO_H3_FFN_CHUNK_TOKENS.get() def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + tokens = hidden_states.shape[-2] if hidden_states.dim() > 1 else 0 + if (self.chunk_tokens and tokens > self.chunk_tokens and not torch.is_grad_enabled() + and not torch.compiler.is_compiling() and hidden_states.numel() == tokens * hidden_states.shape[-1]): + out = torch.empty_like(hidden_states) + for start in range(0, tokens, self.chunk_tokens): + rows = slice(start, start + self.chunk_tokens) + out[..., rows, :] = self._forward(hidden_states[..., rows, :].contiguous()) + return out + return self._forward(hidden_states) + + def _forward(self, hidden_states: torch.Tensor) -> torch.Tensor: if self.use_mxfp8: from fastvideo.layers.mxfp8linear import mxfp8_swiglu_feed_forward @@ -209,6 +237,11 @@ def __init__( prefix=prefix, fa4_packed_varlen=fa4_packed_varlen, ) + # Opt-in inference route: VSA-H3 selection on the block-sparse FP4 + # kernel (see minimax_h3_vsa_fp4); grad and compile keep the generic path. + self._layer_idx = layer_idx_from_prefix(prefix, default=-1) + self._vsa_fp4 = use_vsa and vsa_fp4_requested() + self._vsa_tile_first = use_vsa and envs.FASTVIDEO_H3_VSA_TILE_FIRST.get() self.to_gate_compress: ReplicatedLinear | None = None # None = unchecked; the first forward tests the loaded weight once and # skips the gate branch entirely while it is structurally zero. @@ -255,6 +288,12 @@ def _resolve_gate_compress_for_compile(self) -> None: return if self._gate_compress_active is None: weight = self.to_gate_compress.weight + if weight is None: + # Packed NVFP4 gate: any nonzero E2M1 magnitude (bits 0x7 of + # either nibble) makes the branch live. + packed = self.to_gate_compress._nvfp4_weight + self._gate_compress_active = bool((packed & 0x77).any()) + return # bool() on a DTensor reduction resolves collectively, so every # rank caches the same answer. self._gate_compress_active = bool((weight != 0).any()) @@ -282,9 +321,41 @@ def forward( rotary_emb: tuple[torch.Tensor, torch.Tensor] | None, original_seq_len: int, ) -> torch.Tensor: - query, _ = self.to_q(hidden_states) - key, _ = self.to_k(hidden_states) - value, _ = self.to_v(hidden_states) + if (self._vsa_fp4 and rotary_emb is not None and not torch.is_grad_enabled() + and not torch.compiler.is_compiling()): + meta = get_forward_context().attn_metadata + # Exempt mode lists the first prefix tile for every query, which the + # FP4 kernel relies on to start each row from a finite running max. + if isinstance(meta, MiniMaxH3VSAMetadata) and meta.exempt: + use_fused_rope = self.fuse_qknorm_rope and _can_run_minimax_h3_fusion(hidden_states) + if not model_parallel_is_initialized() or get_sp_world_size() == 1: + hidden_states = vsa_fp4_attention(self, hidden_states, rotary_emb, meta, use_fused_rope) + else: + hidden_states = vsa_fp4_attention_sp(self, hidden_states, rotary_emb, meta, use_fused_rope, + get_sp_group()) + with STAGES.span("out_proj"): + hidden_states, _ = self.to_out(hidden_states) + return hidden_states + if (self._vsa_tile_first and hidden_states.is_cuda and rotary_emb is not None + and not torch.is_grad_enabled() and not torch.compiler.is_compiling() + and (not model_parallel_is_initialized() or get_sp_world_size() == 1)): + meta = get_forward_context().attn_metadata + if isinstance(meta, MiniMaxH3VSAMetadata) and meta.tile_elems == 64: + use_fused_rope = self.fuse_qknorm_rope and _can_run_minimax_h3_fusion(hidden_states) + hidden_states = vsa_tile_first_attention(self, hidden_states, rotary_emb, meta, use_fused_rope) + with STAGES.span("out_proj"): + hidden_states, _ = self.to_out(hidden_states) + return hidden_states + with STAGES.span("qkv_proj"): + # All three projections see the same activations. Reuse their FP8 + # quantization when the loaded methods have identical granularity. + from fastvideo.models.dits.minimax_h3_vsa_fp4 import _shared_input_projections + if not torch.is_grad_enabled() and not torch.compiler.is_compiling(): + query, key, value = _shared_input_projections((self.to_q, self.to_k, self.to_v), hidden_states) + else: + query, _ = self.to_q(hidden_states) + key, _ = self.to_k(hidden_states) + value, _ = self.to_v(hidden_states) 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)) @@ -308,16 +379,18 @@ def forward( gate_compress, _ = self.to_gate_compress(hidden_states) extra_attention_kwargs["gate_compress"] = gate_compress.unflatten( -1, (self.num_attention_heads, self.attention_head_dim)) - hidden_states, _ = self.distributed_attention( - query, - key, - value, - original_seq_len=original_seq_len, - freqs_cis=None, - **extra_attention_kwargs, - ) + with STAGES.span("attention"): + hidden_states, _ = self.distributed_attention( + query, + key, + value, + original_seq_len=original_seq_len, + freqs_cis=None, + **extra_attention_kwargs, + ) hidden_states = hidden_states.flatten(2, 3).type_as(query) - hidden_states, _ = self.to_out(hidden_states) + with STAGES.span("out_proj"): + hidden_states, _ = self.to_out(hidden_states) return hidden_states @@ -421,7 +494,51 @@ def __init__( prefix=f"{prefix}.linear", ) + def enable_host_cache(self, table: dict | None = None) -> None: + """Keep the projection in pinned host memory and cache its output per timestep set. + + The modulation is a pure function of the timestep embedding, and few-step checkpoints sample a fixed + timestep ladder, so each block's output is a small constant table. The weights (the largest bf16 tensors + in the DiT) then never occupy device memory: a cache miss copies them in for one matmul. + """ + weight, bias = self.linear.weight, self.linear.bias + if table is not None and (weight is None or bias is None): + # Table-only load skipped these tensors entirely. + self._host_weight = self._host_bias = None + self._modulation_cache = dict(table) + self._cache_key = None + return + weight, bias = weight.data, bias.data + if table is None: + self._host_weight = weight.to("cpu").pin_memory() + self._host_bias = bias.to("cpu").pin_memory() + else: + # Precomputed modulation for a fixed timestep ladder: the projection weights are not needed at all. + self._host_weight = self._host_bias = None + self.linear.weight.data = torch.empty(0, dtype=weight.dtype) + self.linear.bias.data = torch.empty(0, dtype=bias.dtype) + self._modulation_cache: dict[Any, torch.Tensor] = dict(table or {}) + self._cache_key: Any = None + def forward(self, temb: torch.Tensor) -> tuple[torch.Tensor, ...]: + cache = getattr(self, "_modulation_cache", None) + if cache is not None: + out = cache.get(self._cache_key) + if out is None: + if self._host_weight is None: + raise RuntimeError(f"No precomputed AdaLN modulation for timestep key {self._cache_key}; the " + "table only covers the checkpoint's fixed ladder. Load without " + "FASTVIDEO_H3_ADALN_TABLE to use other timesteps.") + x = F.silu(temb) if self.apply_silu else temb + weight = self._host_weight.to(temb.device, non_blocking=True) + bias = self._host_bias.to(temb.device, non_blocking=True) + out = F.linear(x.to(weight.dtype), weight, bias) + if self._cache_key is not None: + cache[self._cache_key] = out + if envs.FASTVIDEO_H3_ADALN_DUMP.get(): + # Projection inputs, kept only when dumping, for offline low-rank fits. + self.__dict__.setdefault("_modulation_inputs", {})[self._cache_key] = x.detach() + return out.view(-1, 6 * self.hidden_size).chunk(6, dim=-1) if self.apply_silu: temb = F.silu(temb) temb, _ = self.linear(temb.to(self.linear.weight.dtype)) @@ -546,6 +663,9 @@ def forward( 1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices) with nvtx_range("minimax_h3.transformer_block.self_attention"): attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len) + # The attention input is dead now. Keeping it until assignment below + # overlaps three full-width activations during the residual fusion. + del norm_hidden_states if use_modulate_fusion: with nvtx_range("minimax_h3.transformer_block.modulate_fusion"): hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate( @@ -564,8 +684,11 @@ def forward( 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) - with nvtx_range("minimax_h3.transformer_block.feed_forward"): + del attention_output + with nvtx_range("minimax_h3.transformer_block.feed_forward"), STAGES.span("feed_forward"): feed_forward_output = self.ff(norm_hidden_states) + if use_modulate_fusion and not torch.compiler.is_compiling(): + return _gated_residual(hidden_states, gate_mlp, adaln_indices, feed_forward_output) return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output @@ -891,6 +1014,55 @@ def _rotary_for(self, position_ids: torch.Tensor, dtype: torch.dtype) -> tuple[t self._rope_cache = (position_ids, dtype, value) return value + def attach_step_splice(self, late: nn.Module, from_step: int) -> None: + """Hand denoising steps ``from_step`` onward to ``late`` (same architecture, other weights). + + Early DMD steps fix layout and object count, late ones texture and detail, so two checkpoints + can split the trajectory. ``late`` is kept out of this module's children: it is placed, offloaded + and checkpointed on its own. + """ + object.__setattr__(self, "_splice_late", late) + self._splice_from_step = int(from_step) + + def enable_adaln_host_cache(self, table_path: str | None = None) -> None: + """Move every block's AdaLN projection to pinned host memory behind a per-timestep cache. + + With ``table_path`` (written by FASTVIDEO_H3_ADALN_DUMP), the cache is prefilled from precomputed + modulation tables and the projection weights are dropped entirely. + """ + tables = None + if table_path: + import ast + raw = torch.load(table_path, map_location="cpu") + tables = {int(i): {ast.literal_eval(k): v for k, v in blk.items()} for i, blk in raw.items()} + for index, block in enumerate(self.transformer_blocks): + block.adaln_proj.enable_host_cache(None if tables is None else tables[index]) + self._adaln_host_cache = True + self._adaln_dumped_entries = -1 + + def _move_adaln_tables(self, device: torch.device) -> None: + for block in self.transformer_blocks: + cache = block.adaln_proj._modulation_cache + for key, value in cache.items(): + if value.device != device: + cache[key] = value.to(device) + + def _maybe_dump_adaln_tables(self) -> None: + path = envs.FASTVIDEO_H3_ADALN_DUMP.get() + if not path: + return + entries = sum(len(b.adaln_proj._modulation_cache) for b in self.transformer_blocks) + if entries == self._adaln_dumped_entries: + return + self._adaln_dumped_entries = entries + if model_parallel_is_initialized() and get_sp_group().rank_in_group != 0: + return + torch.save({i: {repr(k): v.detach().cpu() for k, v in b.adaln_proj._modulation_cache.items()} + for i, b in enumerate(self.transformer_blocks)}, path) + first = self.transformer_blocks[0].adaln_proj.__dict__.get("_modulation_inputs", {}) + torch.save({repr(k): v.detach().cpu() for k, v in first.items()}, path + ".inputs") + logger.info("Dumped AdaLN modulation tables (%d entries) to %s", entries, path) + def _refined_text(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor: """The prompt embedding is constant across the denoising loop and the refiner blocks are timestep-free, so refine once per generation @@ -930,6 +1102,10 @@ def forward( text_indices: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: """Predict video and audio velocities from one caller-defined packed layout.""" + late = self.__dict__.get("_splice_late") + if late is not None and get_forward_context().current_timestep >= self._splice_from_step: + return late(hidden_states, audio_hidden_states, encoder_hidden_states, timestep, timestep_indices, + token_tags, position_ids, video_indices, audio_indices, text_indices) if position_ids.ndim != 2 or position_ids.shape[-1] != 3: raise ValueError(f"position_ids must have shape (seq_len, 3), got {tuple(position_ids.shape)}.") sequence_length = position_ids.shape[0] @@ -972,10 +1148,20 @@ def forward( local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0) rotary_emb = (rotary_cos, rotary_sin) + if getattr(self, "_adaln_host_cache", False): + # One host read per forward keys every block's modulation cache by the timestep values. + key = (tuple(timestep.reshape(-1).tolist()), tuple(temb.shape), str(temb.dtype)) + for block in self.transformer_blocks: + block.adaln_proj._cache_key = key + self._move_adaln_tables(temb.device) + # 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}"): + if STAGES.enabled: + logger.info("H3_MEMORY_BLOCK %d allocated=%.3f GiB reserved=%.3f GiB", block_index, + torch.cuda.memory_allocated() / 2**30, torch.cuda.memory_reserved() / 2**30) + with nvtx_range(f"minimax_h3.transformer_block.{block_index}"), STAGES.span("block_total"): packed_hidden_states = block( packed_hidden_states, temb, @@ -997,6 +1183,12 @@ def forward( video_output = video_output.index_select(1, video_indices) audio_output = audio_output.index_select(1, audio_indices) + if getattr(self, "_adaln_host_cache", False): + self._maybe_dump_adaln_tables() + if STAGES.enabled: + stages = STAGES.flush() + if not model_parallel_is_initialized() or get_sp_group().rank_in_group == 0: + logger.info("H3_STAGE_MS %s", json.dumps(stages)) return video_output, audio_output diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py new file mode 100644 index 0000000000..3f3a48ef5b --- /dev/null +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -0,0 +1,426 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Inference fast path: VSA-H3 attention on the block-sparse FP4 kernel. + +Opt-in with ``FASTVIDEO_H3_VSA_FP4=1`` (no-grad, single sequence-parallel +rank, ``fastvideo-kernel`` built with ``attn_qat_infer``). The selection is +VSA-H3's own: tile pooling, top-k block mask and the gated compression branch +are unchanged; only the block-sparse attention itself runs on SageAttention3's +FP4 kernel (BF16 Triton otherwise), with 64-token tiles carried by quadrant +masks on the kernel's 128x128 blocks. + +The attention input is gathered into tile order once per block (one +``hidden_size``-wide pass; pad rows stay zero, so q/k/v pad rows are exactly +zero through the bias-free projections, RMSNorm and RoPE). That replaces the +generic path's concat, four tile scatters and three transposes, and lets q, k +and v share one activation quantization. The output returns to packed order +with one gather before ``to_out``. +""" + +from __future__ import annotations + +import math +from typing import Any + +import torch + +import fastvideo.envs as envs + +from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAMetadata, _build_block_mask, _pool_tiles) + +_BLOCK = 128 + + +class _StageTimer: + """Opt-in (``FASTVIDEO_H3_SP_PROFILE=1``) CUDA-event spans summed per stage over one DiT forward. + + Events are only recorded during the forward; ``flush`` synchronizes once + and returns milliseconds per stage, so enabling it costs one sync per step. + """ + + def __init__(self) -> None: + self.enabled = envs.FASTVIDEO_H3_SP_PROFILE.get() + self._spans: list[tuple[str, torch.cuda.Event, torch.cuda.Event]] = [] + + def span(self, name: str): + import contextlib + if not self.enabled: + return contextlib.nullcontext() + return self._record(name) + + def _record(self, name: str): + import contextlib + + @contextlib.contextmanager + def ctx(): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + yield + end.record() + self._spans.append((name, start, end)) + + return ctx() + + def flush(self) -> dict[str, float]: + torch.cuda.synchronize() + totals: dict[str, float] = {} + for name, start, end in self._spans: + totals[name] = totals.get(name, 0.0) + start.elapsed_time(end) + self._spans.clear() + return {k: round(v, 1) for k, v in totals.items()} + + +STAGES = _StageTimer() + +_fp4_api: Any = None + + +def vsa_fp4_requested() -> bool: + return envs.FASTVIDEO_H3_VSA_FP4.get() + + +def _api() -> Any: + global _fp4_api + if _fp4_api is None: + import attn_qat_infer.api as api + _fp4_api = api + return _fp4_api + + +class _TileLayout: + """Per-step tile-order state shared by every block of one forward.""" + + def __init__(self, meta: MiniMaxH3VSAMetadata, rotary_emb: tuple[torch.Tensor, torch.Tensor]) -> None: + self.tile = int(meta.tile_elems) + self.n_tiles = int(meta.variable_block_sizes.numel()) + self.rows = math.ceil(self.n_tiles * self.tile / _BLOCK) * _BLOCK + self.untile = meta.untile_combined_index + self.row_tile = self.untile // self.tile + self.rotary_src = rotary_emb + cos, sin = rotary_emb + self.cos = cos.new_zeros((self.rows, cos.shape[-1])).index_copy_(0, self.untile, cos) + self.sin = sin.new_zeros((self.rows, sin.shape[-1])).index_copy_(0, self.untile, sin) + self._buf: torch.Tensor | None = None + + def gather_in(self, x: torch.Tensor) -> torch.Tensor: + """Packed ``[B, L, C]`` -> tile-ordered ``[B, rows, C]``; pad rows stay zero. + + The buffer is reused across blocks: pad rows are never written and + every valid row is overwritten, and each block consumes it (q/k/v + projections) before the next block refills it. + """ + shape = (x.shape[0], self.rows, x.shape[-1]) + if self._buf is None or self._buf.shape != shape or self._buf.dtype != x.dtype: + self._buf = x.new_zeros(shape) + return self._buf.index_copy_(1, self.untile, x) + + +def _layout_for(meta: MiniMaxH3VSAMetadata, rotary_emb: tuple[torch.Tensor, torch.Tensor]) -> _TileLayout: + layout = getattr(meta, "_h3_fp4_layout", None) + if layout is None or layout.rotary_src[0] is not rotary_emb[0]: + layout = _TileLayout(meta, rotary_emb) + meta._h3_fp4_layout = layout # type: ignore[attr-defined] + return layout + + +def _shared_input_projections(linears: tuple[Any, ...], x: torch.Tensor) -> list[torch.Tensor]: + """Share compatible FP8 preparation or unit-scale NVFP4 activation quantization. + + Calibrated or dynamic NVFP4 activation scales retain independent preparation. + """ + from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod + from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod + + methods = [linear.quant_method for linear in linears] + same_nvfp4 = all(type(m) is NVFP4QuantizeMethod and m.wants_prequantized_input() + and m.uses_unit_activation_scale(linear) + for m, linear in zip(methods, linears, strict=True)) + same_fp8 = all(type(m) is FP8QuantizeMethod and m.granularity == methods[0].granularity for m in methods) + if not (same_nvfp4 or same_fp8): + return [linear(x)[0] for linear in linears] + pre = methods[0].quantize_input(x) + return [m.apply(linear, x, linear.bias, pre_quantized=pre) for m, linear in zip(methods, linears, strict=True)] + + +def vsa_tile_first_attention(attn: Any, hidden_states: torch.Tensor, + rotary_emb: tuple[torch.Tensor, torch.Tensor], + meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> torch.Tensor: + """Single-rank BF16 VSA with one input scatter instead of a Q/K/V/gate stack. + + The existing backend computes the same tile-64 mask, valid-key handling, + fine attention and compression branch. Bias-free projections keep pad + rows zero. This path is inference-only and keeps the checkpoint layout. + """ + layout = _layout_for(meta, rotary_emb) + logical = layout.n_tiles * layout.tile + heads, dim = attn.num_attention_heads, attn.attention_head_dim + with STAGES.span("tile_input"): + x_tiles = layout.gather_in(hidden_states)[:, :logical] + with STAGES.span("qkv_proj"): + query, key, value = (t.unflatten(-1, (heads, dim)) + for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), x_tiles)) + with STAGES.span("qknorm_rope"): + cos, sin = layout.cos[:logical], layout.sin[:logical] + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + query = fused_qknorm_rope(query, attn.norm_q.weight, cos.to(query.dtype), sin.to(query.dtype), attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos.to(key.dtype), sin.to(key.dtype), attn.norm_k.eps) + else: + query = attn._apply_rotary_emb(attn.norm_q(query), (cos, sin)) + key = attn._apply_rotary_emb(attn.norm_k(key), (cos, sin)) + gate = None + if attn.to_gate_compress is not None and attn._gate_active(): + with STAGES.span("gate_proj"): + gate, _ = attn.to_gate_compress(x_tiles) + gate = gate.unflatten(-1, (heads, dim)) + capture_root = envs.FASTVIDEO_H3_CAPTURE_QKV.get() + if capture_root and attn._layer_idx in (0, 20, 41): + from pathlib import Path + root = Path(capture_root) + root.mkdir(parents=True, exist_ok=True) + capture = root / f"layer-{attn._layer_idx}.pt" + if not capture.exists(): + q_pooled = _pool_tiles(query, meta.variable_block_sizes, meta.tile_elems) + k_pooled = _pool_tiles(key, meta.variable_block_sizes, meta.tile_elems) + scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / dim**0.5 + sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity + mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt, + meta.video_tile_spans, meta.span_sparsities) + # Two heads keep the artifact small while retaining all real keys, + # query rows, per-tile selections and partial-tile validity. + torch.save({"q": query[:, :, :2].transpose(1, 2).contiguous().cpu(), + "k": key[:, :, :2].transpose(1, 2).contiguous().cpu(), + "v": value[:, :, :2].transpose(1, 2).contiguous().cpu(), + "mask": mask[:, :2].cpu(), "vbs": meta.variable_block_sizes.cpu(), + "untile": meta.untile_combined_index.cpu()}, capture) + del q_pooled, k_pooled, scores, mask + with STAGES.span("attention"): + out = attn.distributed_attention.attn_impl.forward(query, key, value, gate, meta) + with STAGES.span("untile_output"): + return out.index_select(1, layout.untile).flatten(2, 3) + + +def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], + meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> torch.Tensor: + """Attention core for ``MiniMaxH3Attention``; returns the pre-``to_out`` ``[B, L, H*D]``.""" + api = _api() + layout = _layout_for(meta, rotary_emb) + heads, dim = attn.num_attention_heads, attn.attention_head_dim + with STAGES.span("qkv_proj_rope"): + x_tiles = layout.gather_in(hidden_states) + query, key, value = (t.unflatten(-1, (heads, dim)) + for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), x_tiles)) + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + cos, sin = layout.cos.to(query.dtype), layout.sin.to(query.dtype) + query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) + else: + rope = (layout.cos, layout.sin) + query = attn._apply_rotary_emb(attn.norm_q(query), rope) + key = attn._apply_rotary_emb(attn.norm_k(key), rope) + + sim_fp8 = envs.FASTVIDEO_H3_SIM_SP_FP8.get() + if sim_fp8: + query, key, value = (_fp8_roundtrip(t) for t in (query, key, value)) + + vbs = meta.variable_block_sizes + logical = layout.n_tiles * layout.tile + with STAGES.span("select_mask"): + q_pooled = _pool_tiles(query[:, :logical], vbs, layout.tile) + k_pooled = _pool_tiles(key[:, :logical], vbs, layout.tile) + scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (dim**0.5) + sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity + mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt, meta.video_tile_spans, + meta.span_sparsities) + q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) + with STAGES.span("fp4_attention"): + # Lists come from vsa_tile_mask_to_fp4_blocks and are in range; skip the per-call host sync. + out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad, + validate=False) + with STAGES.span("out_untile"): + out = out.transpose(1, 2).index_select(1, layout.untile) # [B, L, H, D], packed order + if sim_fp8: + out = _fp8_roundtrip(out) + + if attn.to_gate_compress is not None and attn._gate_active(): + with STAGES.span("gate_compress"): + gate, _ = attn.to_gate_compress(hidden_states) + v_pooled = _pool_tiles(value[:, :logical], vbs, layout.tile) + out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled).permute(0, 2, 1, 3).to(out.dtype) + out = out.addcmul_(out_c.index_select(1, layout.row_tile), gate.unflatten(-1, (heads, dim))) + return out.flatten(2, 3) + + +# --------------------------------------------------------------------------- +# Ulysses sequence parallelism: FP8 head/sequence exchange around the FP4 core +# --------------------------------------------------------------------------- +# +# PCIe-only boxes (e.g. 8x RTX PRO 6000) move ~21 GB/s per GPU in an +# all-to-all, so the Ulysses exchange, not compute, bounds multi-GPU latency. +# q/k/v travel as FP8 (one scale per token and head): the attention kernel +# re-quantizes them to FP4 on arrival, so FP8 transport adds error far below +# that floor while halving the bytes. The VSA gate never travels: each rank +# applies it to its own sequence rows after an all-gather of the tiny per-tile +# compression output. The attention output returns as FP8 too; ``to_out`` +# re-quantizes it to FP4. + +_FP8 = torch.float8_e4m3fn +_FP8_MAX = 448.0 + +# Debug: apply the SP path's FP8 rounding on one GPU (q/k/v after RoPE and the +# attention output, one scale per token and head), to separate exchange +# rounding from sharding errors when comparing SP>1 against SP=1. + + +@torch.compile(dynamic=True, fullgraph=True) +def _fp8_roundtrip(x: torch.Tensor) -> torch.Tensor: + """``[..., H, D]`` BF16 -> FP8 with a per-(token, head) scale -> BF16, as the SP exchange does.""" + xf = x.float() + scale = (xf.abs().amax(dim=-1) / _FP8_MAX).clamp_min(1e-12) + return ((xf / scale[..., None]).to(_FP8).float() * scale[..., None]).to(torch.bfloat16) + + +@torch.compile(dynamic=False, fullgraph=True) +def _pack_heads_fp8(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, + world: int) -> tuple[torch.Tensor, torch.Tensor]: + """``[rows, H, D]`` x3 -> per-destination ``[W, 3, rows, H/W, D]`` FP8 payload and scales.""" + x = torch.stack((query, key, value)).float() + scale = (x.abs().amax(dim=-1) / _FP8_MAX).clamp_min(1e-12) + payload = (x / scale[..., None]).to(_FP8) + _, rows, heads, dim = x.shape + payload = payload.view(3, rows, world, heads // world, dim).permute(2, 0, 1, 3, 4).contiguous() + scale = scale.view(3, rows, world, heads // world).permute(2, 0, 1, 3).contiguous() + return payload, scale + + +@torch.compile(dynamic=False, fullgraph=True) +def _unpack_seq_fp8(payload: torch.Tensor, scale: torch.Tensor, seq_len: int) -> torch.Tensor: + """``[W, 3, rows, Hs, D]`` from every source rank -> ``[3, seq_len, Hs, D]`` BF16 in packed order.""" + x = payload.float() * scale[..., None] + world, _, rows, heads, dim = x.shape + return x.permute(1, 0, 2, 3, 4).reshape(3, world * rows, heads, dim)[:, :seq_len].to(torch.bfloat16) + + +@torch.compile(dynamic=False, fullgraph=True) +def _pack_seq_fp8(out_bhsd: torch.Tensor, untile: torch.Tensor, world: int, + rows: int) -> tuple[torch.Tensor, torch.Tensor]: + """Tile-ordered ``[1, Hs, R, D]`` -> packed, padded, per-destination ``[W, rows, Hs, D]`` FP8 + scales.""" + x = out_bhsd[0].transpose(0, 1).index_select(0, untile).float() # [L, Hs, D] + x = torch.nn.functional.pad(x, (0, 0, 0, 0, 0, world * rows - x.shape[0])) + scale = (x.abs().amax(dim=-1) / _FP8_MAX).clamp_min(1e-12) + payload = (x / scale[..., None]).to(_FP8) + return payload.view(world, rows, *payload.shape[1:]), scale.view(world, rows, scale.shape[-1]) + + +@torch.compile(dynamic=False, fullgraph=True) +def _unpack_heads_fp8(payload: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + """``[W, rows, Hs, D]`` from every head group -> ``[rows, W*Hs, D]`` BF16.""" + x = payload.float() * scale[..., None] + world, rows, heads, dim = x.shape + return x.permute(1, 0, 2, 3).reshape(rows, world * heads, dim).to(torch.bfloat16) + + +@torch.compile(dynamic=False, fullgraph=True) +def _apply_gate(out: torch.Tensor, out_c: torch.Tensor, row_tile: torch.Tensor, gate: torch.Tensor) -> torch.Tensor: + """``out + out_c[row_tile] * gate`` for one rank's rows; out/gate ``[rows, H, D]``, out_c ``[n_tiles, H, D]``.""" + return out + out_c.index_select(0, row_tile) * gate + + +def _all_to_all(payload: torch.Tensor, scale: torch.Tensor, group: Any) -> tuple[torch.Tensor, torch.Tensor]: + import torch.distributed as dist + recv = torch.empty_like(payload) + recv_scale = torch.empty_like(scale) + dist.all_to_all_single(recv.view(torch.uint8), payload.view(torch.uint8), group=group) + dist.all_to_all_single(recv_scale, scale, group=group) + return recv, recv_scale + + +class _SPTileLayout: + """Per-step tile-order state for one rank's head subset.""" + + def __init__(self, meta: MiniMaxH3VSAMetadata, rank: int, local_rows: int) -> None: + self.tile = int(meta.tile_elems) + self.n_tiles = int(meta.variable_block_sizes.numel()) + self.rows = math.ceil(self.n_tiles * self.tile / _BLOCK) * _BLOCK + self.seq_len = int(meta.total_seq_length) + self.untile = meta.untile_combined_index + row_tile = self.untile // self.tile + local = torch.arange(rank * local_rows, (rank + 1) * local_rows, device=row_tile.device) + # Rows past the sequence are SP padding; their outputs are discarded. + self.local_row_tile = row_tile[local.clamp_max(self.seq_len - 1)] + self._buf: torch.Tensor | None = None + + def tiles_from(self, qkv: torch.Tensor) -> torch.Tensor: + """``[3, L, Hs, D]`` packed -> ``[3, rows, Hs, D]`` tile order with zero padding (reused buffer).""" + shape = (3, self.rows, *qkv.shape[2:]) + if self._buf is None or self._buf.shape != shape: + self._buf = qkv.new_zeros(shape) + return self._buf.index_copy_(1, self.untile, qkv) + + +def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], + meta: MiniMaxH3VSAMetadata, use_fused_rope: bool, sp_group: Any) -> torch.Tensor: + """Ulysses-SP attention core on local sequence rows ``[1, rows, C]``; returns pre-``to_out`` ``[1, rows, H*D]``.""" + import torch.distributed as dist + + api = _api() + world, rank = sp_group.world_size, sp_group.rank_in_group + heads, dim = attn.num_attention_heads, attn.attention_head_dim + local_rows = hidden_states.shape[1] + layout = getattr(meta, "_h3_fp4_sp_layout", None) + if layout is None: + layout = _SPTileLayout(meta, rank, local_rows) + meta._h3_fp4_sp_layout = layout # type: ignore[attr-defined] + + with STAGES.span("qkv_proj_rope"): + query, key, value = (t.unflatten(-1, (heads, dim)) + for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), hidden_states)) + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + cos, sin = rotary_emb[0].to(query.dtype), rotary_emb[1].to(query.dtype) + query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) + else: + query = attn._apply_rotary_emb(attn.norm_q(query), rotary_emb) + key = attn._apply_rotary_emb(attn.norm_k(key), rotary_emb) + + with STAGES.span("qkv_pack"): + payload, scale = _pack_heads_fp8(query[0], key[0], value[0], world) + with STAGES.span("qkv_all_to_all"): + payload, scale = _all_to_all(payload, scale, sp_group.device_group) + with STAGES.span("qkv_unpack_tile"): + qkv = layout.tiles_from(_unpack_seq_fp8(payload, scale, layout.seq_len)) # [3, R, Hs, D] + q_t, k_t, v_t = qkv[0:1], qkv[1:2], qkv[2:3] + + vbs = meta.variable_block_sizes + logical = layout.n_tiles * layout.tile + with STAGES.span("select_mask"): + scores = torch.matmul(_pool_tiles(q_t[:, :logical], vbs, layout.tile), + _pool_tiles(k_t[:, :logical], vbs, layout.tile).transpose(-2, -1)) / (dim**0.5) + sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity + mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt, meta.video_tile_spans, + meta.span_sparsities) + q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) + with STAGES.span("fp4_attention"): + out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad, + validate=False) + + with STAGES.span("out_pack"): + payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows) + with STAGES.span("out_all_to_all"): + payload, scale = _all_to_all(payload, scale, sp_group.device_group) + with STAGES.span("out_unpack"): + out = _unpack_heads_fp8(payload, scale) # [rows, H, D] + + if attn.to_gate_compress is not None and attn._gate_active(): + with STAGES.span("gate_compress"): + v_pooled = _pool_tiles(v_t[:, :logical], vbs, layout.tile) + out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled)[0].to(out.dtype) # [Hs, n_tiles, D] + gathered = torch.empty((world, *out_c.shape), dtype=out_c.dtype, device=out_c.device) + with STAGES.span("gate_all_gather"): + dist.all_gather_into_tensor(gathered, out_c.contiguous(), group=sp_group.device_group) + with STAGES.span("gate_apply"): + out_c_all = gathered.flatten(0, 1).transpose(0, 1) # [n_tiles, H, D] + gate, _ = attn.to_gate_compress(hidden_states) + out = _apply_gate(out, out_c_all, layout.local_row_tile, gate[0].unflatten(-1, (heads, dim))) + return out.flatten(1, 2).unsqueeze(0) diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 0c96e636cf..27925846d6 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -43,6 +43,8 @@ from torch import nn from torch.nn.parameter import Parameter +import fastvideo.envs as envs + from fastvideo.distributed import get_tp_world_size from fastvideo.layers.linear import LinearBase, LinearMethodBase from fastvideo.layers.quantization.base_config import QuantizationConfig @@ -54,6 +56,9 @@ _require_flashinfer, ) from fastvideo.models.utils import set_weight_attrs +from fastvideo.logger import init_logger + +logger = init_logger(__name__) NVFP4_GROUP_SIZE = 16 NVFP4_SCALE_LAYOUT = "128x4" @@ -182,6 +187,38 @@ def _nvfp4_linear( ) +_E2M1_VALUES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0) + + +def _fp4_gemm_supported(device: torch.device) -> bool: + """FP4 tensor-core GEMMs exist on Blackwell (sm_100 / sm_120) and newer.""" + return torch.cuda.is_available() and torch.cuda.get_device_capability(device)[0] >= 10 + + +def unswizzle_128x4_scales(scale: torch.Tensor, rows: int, cols: int) -> torch.Tensor: + """FlashInfer ``layout_128x4`` block scales -> row-major ``[rows, cols]`` (E4M3 bytes). + + The swizzled buffer holds the padded matrix ``[ceil(rows/128)*128, ceil(cols/4)*4]`` as + ``(row_tile, col_tile, row % 32, (row // 32) % 4, col % 4)``. + """ + pad_rows, pad_cols = -(-rows // 128) * 128, -(-cols // 4) * 4 + tiles = scale.reshape(-1)[:pad_rows * pad_cols].view(pad_rows // 128, pad_cols // 4, 32, 4, 4) + return tiles.permute(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols] + + +def dequantize_serialized_nvfp4(weight_packed: torch.Tensor, weight_scale: torch.Tensor, global_scale: float, + dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + """Packed E2M1 ``[out, in // 2]`` + swizzled E4M3 scales -> dense ``[out, in]`` weight.""" + out_features, in_features = weight_packed.shape[0], weight_packed.shape[1] * 2 + lut = torch.tensor(_E2M1_VALUES, device=weight_packed.device, dtype=torch.float32) + packed = weight_packed.view(torch.uint8) + values = torch.stack((lut[(packed & 0x0F).long()], lut[(packed >> 4).long()]), dim=-1).reshape( + out_features, in_features) + scales = unswizzle_128x4_scales(weight_scale.view(torch.uint8), out_features, in_features // NVFP4_GROUP_SIZE) + scales = scales.view(torch.float8_e4m3fn).float().repeat_interleave(NVFP4_GROUP_SIZE, dim=1) + return (values * scales / global_scale).to(dtype) + + class MiniMaxH3SerializedNVFP4Config(QuantizationConfig): """Serialized 16-group NVFP4 contract for the H3 text encoder. @@ -267,15 +304,20 @@ def validate_runtime(self, device: torch.device) -> None: raise RuntimeError(f"MiniMax-H3 serialized NVFP4 requires a CUDA device; got {device.type!r}") capability = torch.cuda.get_device_capability(device) capability_number = capability[0] * 10 + capability[1] + if get_tp_world_size() > 1: + raise NotImplementedError("MiniMax-H3 serialized NVFP4 supports a single GPU: packed FP4 columns and " + "128x4 swizzled scale rows cannot be narrowed per tensor-parallel rank") if capability_number < self.get_min_capability(): - raise RuntimeError("MiniMax-H3 serialized NVFP4 requires GPU capability " - f"sm{self.get_min_capability()} or newer, got sm{capability_number}") + if capability_number < 80: + raise RuntimeError("MiniMax-H3 serialized NVFP4 needs bf16 compute (sm80+) for its de-quantized " + f"fallback, got sm{capability_number}") + logger.warning( + "MiniMax-H3 serialized NVFP4 on sm%d: no FP4 GEMM, each linear de-quantizes its weight to bf16 " + "per call", capability_number) + return if capability[0] not in (10, 12): raise RuntimeError("MiniMax-H3 serialized NVFP4 runs FlashInfer's Blackwell FP4 GEMM; " f"got unsupported sm{capability_number}") - if get_tp_world_size() > 1: - raise NotImplementedError("MiniMax-H3 serialized NVFP4 supports a single GPU: packed FP4 columns and " - "128x4 swizzled scale rows cannot be narrowed per tensor-parallel rank") sf_layout, _, _ = _require_flashinfer() if not hasattr(sf_layout, "layout_128x4"): raise RuntimeError("The installed flashinfer has no SfLayout.layout_128x4; MiniMax-H3 serialized NVFP4 " @@ -304,6 +346,7 @@ def load(param: torch.Tensor, loaded_weight: torch.Tensor, *args: Any, **kwargs: f"got {loaded_weight.dtype} for a {param.dtype} parameter of shape {tuple(param.shape)}") return base_loader(param, loaded_weight, *args, **kwargs) + load._h3_base_loader = base_loader return load @@ -377,6 +420,11 @@ def process_weights_after_loading(self, layer: nn.Module) -> None: # ``mm_fp4`` folds both global scales into one multiplier. Activations use a # unit global scale, so the multiplier is the inverse weight global scale. device = weight_scale.device + # Serialized weights are immutable between post-load hooks. Keeping the + # validated scalar on the host avoids a CUDA synchronization per linear + # on the BF16 fallback used by consumer GPUs. + layer._nvfp4_dequant_global_scale = global_scale + layer._nvfp4_fused_dequant = envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.get() layer.register_buffer("_nvfp4_alpha", torch.tensor(1.0 / global_scale, dtype=torch.float32, device=device), persistent=False) layer.register_buffer("_nvfp4_x_global_scale", torch.ones((), dtype=torch.float32, device=device), @@ -385,12 +433,27 @@ def process_weights_after_loading(self, layer: nn.Module) -> None: @staticmethod def _apply_finalized(layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor: x = _coerce_fp4_input_dtype(x) + if not _fp4_gemm_supported(layer.weight_packed.device): + # Pre-Blackwell GPUs have no FP4 GEMM: expand this layer's weight to bf16 for the one call. + # The encoder runs once per request, so the transient weight is cheaper than keeping a bf16 copy. + dequantize = dequantize_serialized_nvfp4 + if layer._nvfp4_fused_dequant: + from fastvideo.layers.quantization.nvfp4_dequant import dequantize_nvfp4_cuda + dequantize = dequantize_nvfp4_cuda + weight = dequantize(layer.weight_packed, layer.weight_scale, layer._nvfp4_dequant_global_scale, x.dtype) + return torch.nn.functional.linear(x, weight, None if bias is None else bias.to(x.dtype)) original_shape = x.shape if x.numel() == 0: # An empty prompt has nothing to quantize; the FP4 kernels are not defined for zero rows. return x.new_zeros(*original_shape[:-1], layer.output_size_per_partition, dtype=torch.bfloat16) - x_fp4, x_scale = _quantize_activation_nvfp4(x.reshape(-1, original_shape[-1]), layer._nvfp4_x_global_scale) - output = _nvfp4_linear(x_fp4, x_scale, layer.weight_packed, layer.weight_scale, layer._nvfp4_alpha) + x_global_scale, alpha = layer._nvfp4_x_global_scale, layer._nvfp4_alpha + if alpha.device != layer.weight_packed.device: + # A streamed encoder layer is finalized on the host and only its packed weights visit the GPU; + # FlashInfer requires the scalars on the GEMM's device. + x_global_scale = x_global_scale.to(layer.weight_packed.device, non_blocking=True) + alpha = alpha.to(layer.weight_packed.device, non_blocking=True) + x_fp4, x_scale = _quantize_activation_nvfp4(x.reshape(-1, original_shape[-1]), x_global_scale) + output = _nvfp4_linear(x_fp4, x_scale, layer.weight_packed, layer.weight_scale, alpha) if bias is not None: # The GEMM emits bf16; keep it that way whatever dtype the bias was built in. output = output + bias.to(output.dtype) diff --git a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py index c348140d1d..3f958ec1b1 100644 --- a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py +++ b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py @@ -690,8 +690,17 @@ def encode_ids( if (pixel_values_videos is None) != (video_grid_thw is None): raise ValueError("pixel_values_videos and video_grid_thw must be provided together") + stream_device = getattr(self, "_h3_encoder_layerwise_device", None) + if stream_device is not None and (pixel_values is not None or pixel_values_videos is not None): + raise ValueError("Layerwise H3 encoder currently supports text-only conditioning; " + "disable FASTVIDEO_H3_ENCODER_LAYERWISE for visual references") + input_ids = input_ids.unsqueeze(0) - inputs_embeds = self.language_model.embed_tokens(input_ids) + embedding_ids = input_ids.to("cpu") if stream_device is not None else input_ids + inputs_embeds = self.language_model.embed_tokens(embedding_ids) + if stream_device is not None: + inputs_embeds = inputs_embeds.to(stream_device) + input_ids = input_ids.to(stream_device) image_mask = None video_mask = None @@ -746,6 +755,25 @@ def encode_ids( raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}") return hidden_states[0] + def prepare_layerwise_offload(self, device: torch.device, *, pin_cpu_memory: bool = True) -> None: + """Stream language layers for text-only CUDA inference, retaining embeddings on CPU.""" + if getattr(self, "_h3_encoder_layerwise_device", None) is not None: + return + if device.type != "cuda": + raise ValueError("Layerwise H3 encoder requires CUDA") + from fastvideo.distributed import get_tp_world_size + from fastvideo.hooks.layerwise_offload import enable_layerwise_offload + + if get_tp_world_size() != 1: + raise ValueError("Layerwise H3 encoder requires tensor parallel size 1") + self.to("cpu") + self.language_model.rotary_emb.to(device) + if self.language_model.norm is not None: + self.language_model.norm.to(device) + enable_layerwise_offload(self.language_model, resident_blocks=0, cyclic=False, + pin_cpu_memory=pin_cpu_memory) + self._h3_encoder_layerwise_device = device + def forward( self, input_ids: torch.Tensor, @@ -777,10 +805,34 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}") parameter = parameters[name] loader = getattr(parameter, "weight_loader", default_weight_loader) - loader(parameter, tensor) + base_loader = getattr(loader, "_h3_base_loader", loader) + copy_only = (base_loader is default_weight_loader or getattr(base_loader, "__func__", None) in ( + ColumnParallelLinear.weight_loader, RowParallelLinear.weight_loader, + VocabParallelEmbedding.weight_loader)) + if getattr(base_loader, "__func__", None) is VocabParallelEmbedding.weight_loader: + output_dim = getattr(parameter, "output_dim", None) + copy_only = (not getattr(parameter, "is_gguf_weight_type", False) + and getattr(parameter, "packed_dim", None) is None + and (output_dim is None or (tensor.ndim > output_dim + and tensor.shape[output_dim] == base_loader.__self__.org_vocab_size))) + if (getattr(self, "_h3_checkpoint_backed_cpu", False) and copy_only + and parameter.device.type == tensor.device.type == "cpu" + and parameter.shape == tensor.shape and parameter.dtype == tensor.dtype): + # TP=1 loaders only copy already matching tensors. Keep the mapping + # so the OS can reclaim encoder checkpoint pages during denoising. + # Preserve the Parameter and its loader/quantization attributes. + parameter.data = tensor.detach() + else: + loader(parameter, tensor) loaded.add(name) return loaded + def enable_checkpoint_backed_cpu_load(self) -> None: + """Retain immutable CPU checkpoint storage for single-GPU streamed inference.""" + if get_tp_world_size() != 1: + raise ValueError("Checkpoint-backed H3 encoder requires tensor parallel size 1") + self._h3_checkpoint_backed_cpu = True + 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 diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 148a6d1bb0..fcaed55507 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -13,6 +13,8 @@ import torch import torch.distributed as dist + +import fastvideo.envs as envs import torch.nn as nn from safetensors.torch import load_file as safetensors_load_file, safe_open from torch.distributed import init_device_mesh @@ -420,6 +422,12 @@ def load_model( with target_device: model = model_cls(model_config) # type: ignore + retain_checkpoint = getattr(model, "enable_checkpoint_backed_cpu_load", None) + checkpoint_backed_cpu = (target_device.type == "cpu" and envs.FASTVIDEO_H3_ENCODER_LAYERWISE.get() + and callable(retain_checkpoint) and not fastvideo_args.pin_cpu_memory) + if checkpoint_backed_cpu: + retain_checkpoint() + 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): if os.path.isdir(checkpoint_path): @@ -460,16 +468,34 @@ def load_model( f"checkpoint: {weights_not_loaded}") if checkpoint_quant_config is not None: - processed_linears = _process_quantized_text_encoder_weights(model, runtime_device) + # NVFP4 validation and scalar derivation work on the host. Moving + # packed layers to CUDA and back would discard checkpoint mappings. + process_device = (target_device if checkpoint_backed_cpu and checkpoint_quant_config.get_name() == "nvfp4" + else runtime_device) + processed_linears = _process_quantized_text_encoder_weights(model, process_device) logger.info("Validated %d serialized %s text-encoder linears", processed_linears, checkpoint_quant_config.get_name()) # Explicitly move model to target device after loading weights model = model.to(target_device) + prepare_layerwise = getattr(model, "prepare_layerwise_offload", None) + if envs.FASTVIDEO_H3_ENCODER_LAYERWISE.get() and callable(prepare_layerwise): + if target_device.type != "cpu": + raise ValueError("Layerwise H3 encoder requires text_encoder_cpu_offload=True") + prepare_layerwise(runtime_device, pin_cpu_memory=fastvideo_args.pin_cpu_memory) + use_cpu_offload = False + logger.info("Enabled text-only layerwise H3 encoder with CPU token embeddings") + from fastvideo.platforms import current_platform - if use_cpu_offload: + if use_cpu_offload and checkpoint_quant_config is not None: + logger.info( + "Skipping FSDP CPU offload for serialized %s text encoder; " + "packed uint8 weights are not FSDP-shardable", + checkpoint_quant_config.get_name(), + ) + elif use_cpu_offload: pin_cpu_memory = fastvideo_args.pin_cpu_memory and is_pin_memory_available() # Disable FSDP for MPS as it's not compatible if current_platform.is_mps(): @@ -876,6 +902,14 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # Find all safetensors files safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) + int8_convrot_path = None + if class_name == "AutoencoderKLMiniMaxH3": + from fastvideo.models.vaes.minimax_h3_int8_convrot import ( + dense_vae_safetensors, + find_int8_convrot_vae_path, + ) + int8_convrot_path = find_int8_convrot_vae_path(model_path) + safetensors_list = dense_vae_safetensors(safetensors_list) if not safetensors_list: raise ValueError(f"No safetensors files found in {model_path}") # Common case: a single `.safetensors` checkpoint file. @@ -911,6 +945,9 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # strictly so missing/unexpected keys are surfaced early. strict_load = class_name in {"AutoencoderKL", "AutoencoderKLMiniMaxH3"} vae.load_state_dict(loaded, strict=strict_load) + if class_name == "AutoencoderKLMiniMaxH3" and int8_convrot_path is not None: + from fastvideo.models.vaes.minimax_h3_int8_convrot import overlay_minimax_h3_int8_convrot_decoder + overlay_minimax_h3_int8_convrot_decoder(vae, int8_convrot_path) if (class_name == "AutoencoderKLWan" and getattr(vae.config, "use_light_vae", False) and target_device.type == "cuda" and hasattr(vae, "optimize_memory_format")): vae.optimize_memory_format() @@ -1068,6 +1105,8 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) if not safetensors_list: raise ValueError(f"No safetensors files found in {model_path}") + from fastvideo.layers.quantization.nvfp4_config import dense_transformer_safetensors + safetensors_list = dense_transformer_safetensors(safetensors_list) # arch_config can infer architecture from weight keys (e.g. Flux2 layer counts) update_fn = getattr(dit_config.arch_config, "update_from_weight_keys", None) @@ -1120,6 +1159,16 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): os.environ.get("RANK", "0"), resolved.name if resolved else "automatic selection", local_main_process_only=False) + # Layerwise offload keeps every block's weights in pinned host memory, so load them on the CPU and + # attach the hooks before anything moves to the GPU; loading on the GPU first would need the whole + # DiT resident once, which is exactly what offload exists to avoid on small cards. + layerwise_load = (fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload + and not fastvideo_args.use_fsdp_inference) + # The AdaLN host cache also needs the projection weights to stay off the device from the start. + adaln_table = envs.FASTVIDEO_H3_ADALN_TABLE.get() or None + adaln_host_cache = (fastvideo_args.inference_mode and not fastvideo_args.use_fsdp_inference + and (envs.FASTVIDEO_H3_ADALN_CACHE.get() or adaln_table is not None)) + layerwise_load = layerwise_load or adaln_host_cache model = maybe_load_fsdp_model( model_cls=model_cls, init_params={ @@ -1127,7 +1176,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): "hf_config": hf_config }, weight_dir_list=safetensors_list, - device=get_local_torch_device(), + device=torch.device("cpu") if layerwise_load else get_local_torch_device(), hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim, hsdp_shard_dim=fastvideo_args.hsdp_shard_dim, strict=strict_load, @@ -1161,15 +1210,40 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): model = model.eval() + if adaln_host_cache and hasattr(model, "enable_adaln_host_cache"): + model.enable_adaln_host_cache(adaln_table) + logger.info("AdaLN modulation: %s", "precomputed tables from " + adaln_table if adaln_table + else "projections in pinned host memory behind a per-timestep cache") + if layerwise_load and not fastvideo_args.dit_layerwise_offload: + model = model.to(get_local_torch_device()) + if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload: # Check if model has nn.ModuleList for layerwise offload compatibility has_module_list = any(isinstance(m, nn.ModuleList) for m in model.children()) if has_module_list: - enable_layerwise_offload(model) + enable_layerwise_offload(model, pin_cpu_memory=fastvideo_args.pin_cpu_memory) + # Blocks now hold placeholders; the remaining (non-block) weights and buffers belong on the GPU. + model = model.to(get_local_torch_device()) else: logger.warning( "Layerwise offload requested but model %s does not have " "nn.ModuleList structure. Skipping layerwise offload.", cls_name) + if layerwise_load: + model = model.to(get_local_torch_device()) + # FASTVIDEO_H3_SPLICE_TRANSFORMER=: a second checkpoint of the same architecture + # runs denoising steps FASTVIDEO_H3_SPLICE_FROM_STEP (default 4) onward. + # Only the primary ``transformer`` component splices; the spliced load itself never does. + splice_path = envs.FASTVIDEO_H3_SPLICE_TRANSFORMER.get() + if (splice_path and not getattr(self, "_loading_splice", False) and hasattr(model, "attach_step_splice") + and os.path.basename(os.path.normpath(model_path)) == "transformer"): + self._loading_splice = True + try: + late = self.load(splice_path, fastvideo_args) + finally: + self._loading_splice = False + from_step = envs.FASTVIDEO_H3_SPLICE_FROM_STEP.get() + model.attach_step_splice(late, from_step) + logger.info("Step splice: steps >= %d run the transformer from %s", from_step, splice_path) return model diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 956850a508..87b53cf653 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -36,6 +36,54 @@ def _summarize_param_names(names: set[str]) -> str: return ", ".join(f"{family} x{count}" if count > 1 else family for family, count in sorted(families.items())) +def _prequantized_fp8_prefixes(weight_files: list[str]) -> list[str]: + """Checkpoint prefixes stored as FP8 W8A8: float8_e4m3fn ``weight`` plus per-channel float32 ``weight_scale``.""" + from safetensors import safe_open + + fp8_weights: set[str] = set() + scales: set[str] = set() + for path in weight_files: + with safe_open(path, framework="pt", device="cpu") as reader: + for key in reader.keys(): # noqa: SIM118 + if key.endswith(".weight") and reader.get_slice(key).get_dtype() == "F8_E4M3": + fp8_weights.add(key[:-len(".weight")]) + elif key.endswith(".weight_scale") and reader.get_slice(key).get_dtype() == "F32": + scales.add(key[:-len(".weight_scale")]) + if fp8_weights - scales: + raise ValueError(f"FP8 checkpoint weights without a weight_scale: {sorted(fp8_weights - scales)[:4]}") + return sorted(fp8_weights) + + +def _load_prequantized_fp8(model: nn.Module, weight_files: list[str], targets: dict[str, str], + device: torch.device) -> None: + """Attach FP8 W8A8 checkpoint tensors to their FP8 linears as ``convert_model_to_fp8`` would (per channel).""" + from safetensors import safe_open + + from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod + + modules = dict(model.named_modules()) + wanted = {f"{prefix}.{suffix}" for prefix in targets for suffix in ("weight", "weight_scale")} + tensors: dict[str, torch.Tensor] = {} + for path in weight_files: + with safe_open(path, framework="pt", device="cpu") as reader: + for key in wanted.intersection(reader.keys()): + tensors[key] = reader.get_tensor(key) + for prefix, target in targets.items(): + module = modules.get(target[:-len(".weight")]) + quant_method = getattr(module, "quant_method", None) + if not isinstance(quant_method, FP8QuantizeMethod): + raise RuntimeError(f"Pre-quantized FP8 weight {prefix!r} maps to {target!r}, which is not an FP8 linear; " + "load the checkpoint with an FP8 quant_config covering it.") + # The checkpoint scales are per output channel; activations then use per-token scales. + quant_method.granularity = "channel" + module.register_buffer("_fp8_weight", tensors[f"{prefix}.weight"].to(device).contiguous(), persistent=False) + module.register_buffer("_fp8_weight_scale", + tensors[f"{prefix}.weight_scale"].reshape(-1).to(device=device, dtype=torch.float32), + persistent=False) + module._parameters.pop("weight", None) + logger.info("Loaded %d pre-quantized FP8 linears", len(targets)) + + def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lora_merge: bool = False) -> None: """Quantize inference linear weights after checkpoint loading. @@ -75,17 +123,28 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor convert_model_to_mxfp8, ) + # NVFP4 may share a model with FP8 linears (e.g. NVFP4 FFN + FP8 attention); handle that pair before the + # per-module walk so the result does not depend on which quantized module comes first. + nvfp4_modules = [m for m in model.modules() if isinstance(getattr(m, "quant_method", None), NVFP4QuantizeMethod)] + if nvfp4_modules: + mixed_fp8 = any(isinstance(getattr(m, "quant_method", None), FP8QuantizeMethod) for m in model.modules()) + if any(getattr(module, "_nvfp4_weight", None) is not None for module in nvfp4_modules): + logger.info("NVFP4 packed export already populated; skipping runtime weight conversion") + elif defer_weight_conversion_until_lora_merge: + logger.info("Deferring NVFP4 weight conversion until the inference LoRA merge completes") + return + else: + logger.info("Converting loaded model weights for NVFP4 linear layers") + convert_model_to_nvfp4(model) + if mixed_fp8: + logger.info("Converting the FP8 linears of a mixed NVFP4/FP8 model") + convert_model_to_fp8(model) + return + qat_train_attached = 0 qat_train_skipped = 0 for mod in model.modules(): qm = getattr(mod, "quant_method", None) - if isinstance(qm, NVFP4QuantizeMethod): - if defer_weight_conversion_until_lora_merge: - logger.info("Deferring NVFP4 weight conversion until the inference LoRA merge completes") - return - logger.info("Converting loaded model weights for NVFP4 linear layers") - convert_model_to_nvfp4(model) - return if isinstance(qm, NVFP4QATQuantizeMethod): logger.info("Converting loaded model weights for NVFP4-QAT linear layers") convert_model_to_fp4(model) @@ -214,6 +273,22 @@ def maybe_load_fsdp_model( """ _validate_fsdp_inference_quantization(init_params, fsdp_inference) + from fastvideo.layers.quantization.nvfp4_config import ( + dense_transformer_safetensors, + find_minimax_h3_nvfp4_dit_export, + load_minimax_h3_nvfp4_dit_export, + nvfp4_linear_weight_param_names, + ) + packed_candidate = find_minimax_h3_nvfp4_dit_export(weight_dir_list) + weight_dir_list = dense_transformer_safetensors(weight_dir_list) + quant_config = getattr(init_params.get("config"), "quant_config", None) + packed_profiles = ("h3_dit", "h3_dit_ffn", "h3_dit_vsa") + packed_nvfp4_export = (packed_candidate + if getattr(quant_config, "layer_profile", None) in packed_profiles else None) + if packed_nvfp4_export is not None and lora_path is not None: + raise ValueError("Packed MiniMax-H3 NVFP4 DiT export cannot be combined with lora_path; " + "merge the adapter before exporting, or load without the packed file.") + # NOTE(will): cast_forward_inputs=True shouldn't be needed as we are # manually casting the inputs to the model mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=False) @@ -275,11 +350,21 @@ def maybe_load_fsdp_model( fsdp_shard_conditions=model._fsdp_shard_conditions, pin_cpu_memory=pin_cpu_memory) - # Host offload is already disabled on unified memory (GB10). Staging the - # 35B FastH3 DiT on CPU and then copying to CUDA doubled that working set - # and took minutes. Follow cpu_offload: read onto the accelerator. - weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=cpu_offload) - logger.info("Loading transformer weights with to_cpu=%s", cpu_offload) + nvfp4_skip_param_names: set[str] = set() + if packed_nvfp4_export is not None: + nvfp4_skip_param_names = nvfp4_linear_weight_param_names(model) + if not nvfp4_skip_param_names: + logger.warning( + "Found %s next to the transformer shards but no NVFP4 linears; " + "ignoring the packed export. Set NVFP4Config(layer_profile='h3_dit').", + packed_nvfp4_export, + ) + packed_nvfp4_export = None + # CPU-targeted layerwise/table loads must not stage the full checkpoint on + # the GPU before offload hooks or skipped projection weights are applied. + load_weights_to_cpu = device.type == "cpu" or cpu_offload or packed_nvfp4_export is not None + weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=load_weights_to_cpu) + logger.info("Loading transformer weights with to_cpu=%s", load_weights_to_cpu) param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping) dense_lora_patch = DenseLoRAPatch.from_adapter( lora_path, @@ -287,9 +372,6 @@ def maybe_load_fsdp_model( strength=lora_strength, ) if dense_lora_patch is not None: - # H3's compression gate is created only by the VSA attention backend. Loading a - # VSA student under dense attention would otherwise warn about 50 unmatched - # replacements and continue with a silently incomplete model. model_parameter_names = {name for name, _ in model.named_parameters()} missing_vsa_gates = sorted(name for name in dense_lora_patch.replacement_parameters if "gate_compress" in name and name not in model_parameter_names) @@ -298,6 +380,20 @@ def maybe_load_fsdp_model( "This LoRA adapter provides MiniMax H3 VSA compression gates, but the selected attention backend " "did not construct them. Use attention_backend='VIDEO_SPARSE_ATTN_H3'. Missing parameters: " + ", ".join(missing_vsa_gates[:3]) + (" ..." if len(missing_vsa_gates) > 3 else "")) + skip_param_names = set(nvfp4_skip_param_names) + fp8_targets: dict[str, str] = {} + fp8_prefixes = _prequantized_fp8_prefixes(weight_dir_list) + if fp8_prefixes: + fp8_targets = {prefix: param_names_mapping_fn(f"{prefix}.weight")[0] for prefix in fp8_prefixes} + skip_param_names |= set(fp8_targets.values()) + fp8_keys = {f"{prefix}.{suffix}" for prefix in fp8_prefixes for suffix in ("weight", "weight_scale")} + weight_iterator = ((name, tensor) for name, tensor in weight_iterator if name not in fp8_keys) + if envs.FASTVIDEO_H3_ADALN_TABLE.get(): + # Precomputed AdaLN modulation replaces the per-block projections; never read their weights. + skip_param_names |= {name for name, _ in model.named_parameters() + if re.fullmatch(r"transformer_blocks\.\d+\.adaln_proj\.linear\.(weight|bias)", name)} + logger.info("Skipping %d AdaLN projection tensors (precomputed modulation tables)", + sum("adaln_proj" in n for n in skip_param_names)) load_model_from_full_model_state_dict( model, weight_iterator, @@ -307,7 +403,12 @@ def maybe_load_fsdp_model( cpu_offload=cpu_offload, param_names_mapping=param_names_mapping_fn, dense_lora_patch=dense_lora_patch, + skip_param_names=skip_param_names or None, ) + if packed_nvfp4_export is not None: + load_minimax_h3_nvfp4_dit_export(model, packed_nvfp4_export, device=device) + if fp8_targets: + _load_prequantized_fp8(model, weight_dir_list, fp8_targets, device=device) if hasattr(model, "materialize_non_persistent_buffers"): model.materialize_non_persistent_buffers(device=device, dtype=default_dtype) for n, p in chain(model.named_parameters(), model.named_buffers()): @@ -585,6 +686,12 @@ def shard_model( fully_shard(model, **root_kwargs) +def _drop_state_dict_parameter(model: nn.Module, param_name: str) -> None: + module_name, _, attr_name = param_name.rpartition(".") + module = model.get_submodule(module_name) if module_name else model + module.register_parameter(attr_name, None) + + # TODO(PY): device mesh for cfg parallel def load_model_from_full_model_state_dict( model: FSDPModule | torch.nn.Module, @@ -596,6 +703,7 @@ def load_model_from_full_model_state_dict( param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None, training_mode: bool = True, dense_lora_patch: DenseLoRAPatch | None = None, + skip_param_names: set[str] | None = None, ) -> _IncompatibleKeys: """ Converting full state dict into a sharded state dict @@ -634,6 +742,11 @@ def load_model_from_full_model_state_dict( # set. for target_param_name in list(custom_param_sd): full_tensor = custom_param_sd.pop(target_param_name) + if skip_param_names and target_param_name in skip_param_names: + continue + if "::" in target_param_name: + logger.warning("Skipping packed NVFP4 export key mixed into dense shards: %s", target_param_name) + continue meta_sharded_param = meta_sd.get(target_param_name) if meta_sharded_param is None: # Some checkpoints include extra entries that are not part of the @@ -701,6 +814,13 @@ def load_model_from_full_model_state_dict( model.reverse_param_names_mapping = reverse_param_names_mapping unused_keys = set(meta_sd.keys()) - set(sharded_sd.keys()) + skipped_unused = unused_keys & skip_param_names if skip_param_names else set() + for skipped_name in skipped_unused: + _drop_state_dict_parameter(model, skipped_name) + unused_keys -= skipped_unused + if skipped_unused: + logger.info("Deferred %d NVFP4 linear weights to the packed DiT export (%s)", + len(skipped_unused), _summarize_param_names(skipped_unused)) if unused_keys: # Say which of these the adapter is about to fill in. Reporting all of them as # "unloaded" was accurate when zero-init was the only outcome; with an adapter diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py new file mode 100644 index 0000000000..d9a12a01f1 --- /dev/null +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -0,0 +1,353 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Comfy ``int8_tensorwise`` + ConvRot overlay for the MiniMax-H3 video VAE decoder. + +The export stores decoder transformer linears as signed int8 with per-output +channel scales and a JSON ``comfy_quant`` marker: + + {"format": "int8_tensorwise", "convrot": true, "convrot_groupsize": 256} + +Weights were rotated offline by a normalized regular Hadamard (group 256). +Inference rotates activations with the same matrix, row-quantizes them, then +runs int8 GEMM. Encoder convolutions stay dense; only the ViT decoder blocks +are quantized. + +Comfy names (``to_qkv``, ``ff.w1`` / ``ff.w2``, ``x_embedder``) are remapped +onto FastVideo's split Q/K/V and ``ff.net`` surface. +""" + +from __future__ import annotations + +import json +import math +import os +from pathlib import Path +from typing import Any + +import torch +import torch.nn as nn +import torch.nn.functional as F +from safetensors.torch import load_file as safetensors_load_file + +import fastvideo.envs as envs + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) + +INT8_CONVROT_FILENAME = "minimax_h3_video_vae_int8_convrot.safetensors" +_COM_FY_QUANT_SUFFIX = ".comfy_quant" +_DEFAULT_GROUP_SIZE = 256 + +_HADAMARD_CACHE: dict[tuple[int, str, torch.dtype], torch.Tensor] = {} + + +def regular_hadamard(size: int, *, device: torch.device, dtype: torch.dtype) -> torch.Tensor: + """Normalized regular Hadamard of order ``4**k`` (ConvRot Theorem 3.3).""" + cache_key = (size, str(device), dtype) + cached = _HADAMARD_CACHE.get(cache_key) + if cached is not None: + return cached + if size < 4 or (size & (size - 1)) != 0 or math.log(size, 4) % 1: + raise ValueError(f"Regular Hadamard size must be a power of 4, got {size}") + h4 = torch.tensor( + [[1, 1, 1, -1], [1, 1, -1, 1], [1, -1, 1, 1], [-1, 1, 1, 1]], + dtype=dtype, + device=device, + ) + hadamard = h4 + current = 4 + while current < size: + hadamard = torch.kron(hadamard, h4) + current *= 4 + hadamard = hadamard / math.sqrt(size) + _HADAMARD_CACHE[cache_key] = hadamard + return hadamard + + +def rotate_activation(x: torch.Tensor, group_size: int) -> torch.Tensor: + features = x.shape[-1] + if features % group_size: + raise ValueError(f"features {features} are not divisible by convrot group_size {group_size}") + groups = features // group_size + hadamard = regular_hadamard(group_size, device=x.device, dtype=x.dtype) + grouped = x.reshape(*x.shape[:-1], groups, group_size) + return torch.matmul(grouped, hadamard).reshape(x.shape) + + +def parse_comfy_quant_marker(blob: torch.Tensor) -> dict[str, Any]: + """Decode the uint8 JSON marker Comfy stores next to each quantized linear.""" + raw = bytes(blob.detach().cpu().contiguous().view(torch.uint8).numpy()) + raw = raw.split(b"\x00", 1)[0] + marker = json.loads(raw.decode("utf-8")) + if not isinstance(marker, dict): + raise ValueError("comfy_quant marker must be a JSON object") + return marker + + +class Int8ConvRotLinear(nn.Module): + """W8A8 linear matching Comfy ``int8_tensorwise`` (+ optional ConvRot).""" + + def __init__( + self, + in_features: int, + out_features: int, + *, + bias: bool, + convrot: bool, + group_size: int, + ) -> None: + super().__init__() + if convrot and in_features % group_size != 0: + raise ValueError( + f"ConvRot weights were rotated with group_size {group_size}, " + f"but in_features {in_features} is not divisible by that group" + ) + self.in_features = in_features + self.out_features = out_features + self.convrot = convrot + self.group_size = group_size + self._transpose_view = envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.get() + self._fused_dequant = envs.FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT.get() + self.register_buffer("weight", torch.empty(out_features, in_features, dtype=torch.int8)) + self.register_buffer("weight_scale", torch.empty(out_features, 1, dtype=torch.float32)) + if bias: + self.bias = nn.Parameter(torch.zeros(out_features), requires_grad=False) + else: + self.register_parameter("bias", None) + + def _dequant_weight(self, dtype: torch.dtype) -> torch.Tensor: + return self.weight.to(dtype) * self.weight_scale.to(dtype) + + @staticmethod + def _dequant_int8_gemm( + acc: torch.Tensor, + x_scale: torch.Tensor, + weight_scale: torch.Tensor, + ) -> torch.Tensor: + # int32 acc is ~K·127² and overflows fp16 before 1/127 scales land. + return acc.float() * x_scale.float() * weight_scale.t().float() + + def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + x_2d = x.reshape(-1, x.shape[-1]).contiguous() + if self.convrot: + x_2d = rotate_activation(x_2d, self.group_size) + row_max = x_2d.abs().amax(dim=-1, keepdim=True).clamp_min(1e-30) + x_scale = row_max / 127.0 + x_q = (x_2d / x_scale).round().clamp(-128, 127).to(torch.int8) + rows = x_q.shape[0] + if rows <= 16: + pad = 17 - rows + x_q = F.pad(x_q, (0, 0, 0, pad)) + x_scale = F.pad(x_scale, (0, 0, 0, pad)) + return x_q, x_scale + + def forward_quantized(self, x_q: torch.Tensor, x_scale: torch.Tensor, + original_shape: tuple[int, ...], dtype: torch.dtype) -> torch.Tensor: + rows = math.prod(original_shape[:-1]) + weight = self.weight.t() + if not self._transpose_view: + weight = weight.contiguous() + acc = torch._int_mm(x_q, weight)[:rows] + if self._fused_dequant and not torch.is_grad_enabled() and not torch.compiler.is_compiling(): + from fastvideo.models.vaes.minimax_h3_int8_kernels import fused_int8_dequant_bias + return fused_int8_dequant_bias(acc, x_scale[:rows], self.weight_scale, self.bias, dtype).view( + *original_shape[:-1], self.out_features) + out = self._dequant_int8_gemm(acc, x_scale[:rows], self.weight_scale) + if self.bias is not None: + out = out + self.bias.float() + return out.to(dtype=dtype).view(*original_shape[:-1], self.out_features) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + original_shape = x.shape + if x.device.type == "cuda" and x.shape[-1] % 8 == 0: + return self.forward_quantized(*self.quantize_input(x), original_shape, x.dtype) + else: + x_2d = x.reshape(-1, original_shape[-1]).contiguous() + if self.convrot: + x_2d = rotate_activation(x_2d, self.group_size) + out = F.linear(x_2d.float(), self._dequant_weight(torch.float32)) + if self.bias is not None: + out = out + self.bias.float() + return out.to(dtype=x.dtype).view(*original_shape[:-1], self.out_features) + + +def shared_int8_projections(layers: tuple[nn.Module, ...], x: torch.Tensor) -> tuple[torch.Tensor, ...]: + """Reuse identical ConvRot/row quantization while retaining each projection's INT8 GEMM.""" + first = layers[0] + compatible = (x.is_cuda and x.shape[-1] % 8 == 0 + and all(isinstance(layer, Int8ConvRotLinear) for layer in layers)) + if compatible: + compatible = all((layer.in_features, layer.convrot, layer.group_size) + == (first.in_features, first.convrot, first.group_size) for layer in layers) + if not compatible: + return tuple(layer(x) for layer in layers) + x_q, x_scale = first.quantize_input(x) + return tuple(layer.forward_quantized(x_q, x_scale, x.shape, x.dtype) for layer in layers) + + +def _int8_linear_from_tensors( + weight: torch.Tensor, + scale: torch.Tensor, + bias: torch.Tensor | None, + marker: dict[str, Any], +) -> Int8ConvRotLinear: + if weight.dtype != torch.int8: + raise ValueError(f"expected int8 weight, got {weight.dtype}") + if scale.ndim == 1: + scale = scale.unsqueeze(1) + convrot = bool(marker.get("convrot", False)) + group_size = int(marker.get("convrot_groupsize", _DEFAULT_GROUP_SIZE)) if convrot else _DEFAULT_GROUP_SIZE + if marker.get("format") not in (None, "int8_tensorwise"): + raise ValueError(f"unsupported comfy_quant format {marker.get('format')!r}") + layer = Int8ConvRotLinear( + weight.shape[1], + weight.shape[0], + bias=bias is not None, + convrot=convrot, + group_size=group_size, + ) + layer.weight.copy_(weight) + layer.weight_scale.copy_(scale.to(torch.float32)) + if bias is not None and layer.bias is not None: + layer.bias.data.copy_(bias) + return layer + + +def _split_qkv( + weight: torch.Tensor, + scale: torch.Tensor, + bias: torch.Tensor | None, + marker: dict[str, Any], + *, + heads: int, + dim_head: int, +) -> tuple[Int8ConvRotLinear, Int8ConvRotLinear, Int8ConvRotLinear]: + """Undo Comfy fused QKV: per-head ``[q, k, v]`` then stacked heads.""" + expected = heads * 3 * dim_head + if weight.shape[0] != expected: + raise ValueError(f"fused to_qkv out_features {weight.shape[0]} != heads*3*dim_head {expected}") + if scale.ndim == 1: + scale = scale.unsqueeze(1) + in_features = weight.shape[1] + weight = weight.view(heads, 3, dim_head, in_features) + scale = scale.view(heads, 3, dim_head, 1) + bias_view = None if bias is None else bias.view(heads, 3, dim_head) + pieces = [] + for index in range(3): + w = weight[:, index].reshape(heads * dim_head, in_features) + s = scale[:, index].reshape(heads * dim_head, 1) + b = None if bias_view is None else bias_view[:, index].reshape(heads * dim_head) + pieces.append(_int8_linear_from_tensors(w.contiguous(), s.contiguous(), b, marker)) + return pieces[0], pieces[1], pieces[2] + + +def _swap_swiglu_halves(tensor: torch.Tensor) -> torch.Tensor: + """Comfy ``[gate; value]`` → FastVideo ``[value; gate]`` along out_features.""" + if tensor.shape[0] % 2: + raise ValueError(f"SwiGLU packed rows {tensor.shape[0]} are not even") + half = tensor.shape[0] // 2 + return torch.cat([tensor[half:], tensor[:half]], dim=0) + + +def overlay_minimax_h3_int8_convrot_decoder(vae: nn.Module, checkpoint_path: str | Path) -> int: + """Swap H3 VAE decoder transformer linears for Comfy int8-convrot weights. + + Dense decoder tensors that the export still stores in float (embed, norms, + scales, ``proj_out``) are copied onto the matching FastVideo modules. + Returns the number of quantized linears installed. + """ + path = Path(checkpoint_path) + tensors = safetensors_load_file(str(path)) + decoder = getattr(vae, "decoder", None) + if decoder is None: + raise ValueError("MiniMax-H3 VAE overlay expected a `.decoder` module") + + if "decoder.x_embedder.weight" in tensors and hasattr(decoder, "proj_in"): + decoder.proj_in.weight.data.copy_(tensors["decoder.x_embedder.weight"]) + if decoder.proj_in.bias is not None and "decoder.x_embedder.bias" in tensors: + decoder.proj_in.bias.data.copy_(tensors["decoder.x_embedder.bias"]) + for name in ( + "decoder.register_tokens", + "decoder.norm_out.weight", + "decoder.norm_out.bias", + "decoder.proj_out.weight", + "decoder.proj_out.bias", + ): + if name not in tensors: + continue + module_name, _, param_name = name.removeprefix("decoder.").rpartition(".") + target = decoder if not module_name else decoder.get_submodule(module_name) + getattr(target, param_name).data.copy_(tensors[name]) + + quantized = 0 + blocks = decoder.transformer_blocks + for index, block in enumerate(blocks): + prefix = f"decoder.transformer_blocks.{index}" + marker = parse_comfy_quant_marker(tensors[f"{prefix}.attn.to_qkv{_COM_FY_QUANT_SUFFIX}"]) + to_q, to_k, to_v = _split_qkv( + tensors[f"{prefix}.attn.to_qkv.weight"], + tensors[f"{prefix}.attn.to_qkv.weight_scale"], + tensors.get(f"{prefix}.attn.to_qkv.bias"), + marker, + heads=int(block.attn.heads), + dim_head=int(block.attn.dim_head), + ) + block.attn.to_q = to_q + block.attn.to_k = to_k + block.attn.to_v = to_v + quantized += 3 + out_marker = parse_comfy_quant_marker(tensors[f"{prefix}.attn.to_out{_COM_FY_QUANT_SUFFIX}"]) + to_out = _int8_linear_from_tensors( + tensors[f"{prefix}.attn.to_out.weight"], + tensors[f"{prefix}.attn.to_out.weight_scale"], + tensors.get(f"{prefix}.attn.to_out.bias"), + out_marker, + ) + block.attn.to_out[0] = to_out + quantized += 1 + w1_marker = parse_comfy_quant_marker(tensors[f"{prefix}.ff.w1{_COM_FY_QUANT_SUFFIX}"]) + w1_weight = _swap_swiglu_halves(tensors[f"{prefix}.ff.w1.weight"]) + w1_scale = _swap_swiglu_halves(tensors[f"{prefix}.ff.w1.weight_scale"]) + w1_bias = tensors.get(f"{prefix}.ff.w1.bias") + if w1_bias is not None: + w1_bias = _swap_swiglu_halves(w1_bias) + block.ff.net[0].proj = _int8_linear_from_tensors(w1_weight, w1_scale, w1_bias, w1_marker) + quantized += 1 + w2_marker = parse_comfy_quant_marker(tensors[f"{prefix}.ff.w2{_COM_FY_QUANT_SUFFIX}"]) + block.ff.net[2] = _int8_linear_from_tensors( + tensors[f"{prefix}.ff.w2.weight"], + tensors[f"{prefix}.ff.w2.weight_scale"], + tensors.get(f"{prefix}.ff.w2.bias"), + w2_marker, + ) + quantized += 1 + for dense_name in ("norm1.weight", "norm2.weight", "scale1", "scale2"): + key = f"{prefix}.{dense_name}" + if key not in tensors: + continue + if dense_name.endswith(".weight"): + module = block.get_submodule(dense_name.rsplit(".", 1)[0]) + module.weight.data.copy_(tensors[key]) + else: + getattr(block, dense_name).data.copy_(tensors[key]) + + device = next(vae.parameters()).device + vae.decoder.to(device) + logger.info("Overlaid MiniMax-H3 int8-convrot VAE decoder: %s linears from %s", quantized, path) + return quantized + + +def is_int8_convrot_vae_path(path: str | Path) -> bool: + raw = os.path.basename(str(path)) + real = os.path.basename(os.path.realpath(str(path))) + return INT8_CONVROT_FILENAME in (raw, real) or "int8_convrot" in raw or "int8_convrot" in real + + +def dense_vae_safetensors(paths: list[str]) -> list[str]: + """Drop the ConvRot overlay so it is not loaded as a dense VAE shard.""" + return [path for path in paths if not is_int8_convrot_vae_path(path)] + + +def find_int8_convrot_vae_path(model_path: str | Path) -> Path | None: + candidate = Path(model_path) / INT8_CONVROT_FILENAME + return candidate if candidate.is_file() else None diff --git a/fastvideo/models/vaes/minimax_h3_int8_kernels.py b/fastvideo/models/vaes/minimax_h3_int8_kernels.py new file mode 100644 index 0000000000..e1e5a9fad4 --- /dev/null +++ b/fastvideo/models/vaes/minimax_h3_int8_kernels.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Eager INT8 VAE epilogue with the reference's separate FP32 operations.""" +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _dequant_bias(Acc, XScale, WScale, Bias, Out, M: tl.constexpr, N: tl.constexpr, + XS: tl.constexpr, WS: tl.constexpr, HAS_BIAS: tl.constexpr, BLOCK: tl.constexpr): + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + rows, cols = offsets // N, offsets % N + valid = rows < M + acc = tl.load(Acc + offsets, valid, 0).to(tl.float32) + x_scale = tl.load(XScale + rows * XS, valid, 0).to(tl.float32) + weight_scale = tl.load(WScale + cols * WS).to(tl.float32) + out = acc * x_scale + out = out * weight_scale + if HAS_BIAS: + out = out + tl.load(Bias + cols).to(tl.float32) + tl.store(Out + offsets, out, valid) + + +def fused_int8_dequant_bias(acc: torch.Tensor, x_scale: torch.Tensor, + weight_scale: torch.Tensor, bias: torch.Tensor | None, + dtype: torch.dtype) -> torch.Tensor: + """Avoid full-size FP32 scaling intermediates; retain both rounding steps.""" + rows, cols = acc.shape + if not acc.is_cuda or acc.dtype != torch.int32 or not acc.is_contiguous(): + raise ValueError("INT8 VAE epilogue requires a contiguous CUDA INT32 matrix") + if x_scale.shape != (rows, 1) or weight_scale.shape != (cols, 1): + raise ValueError("INT8 VAE epilogue requires per-row and per-output-channel scales") + if any(t.device != acc.device for t in (x_scale, weight_scale)): + raise ValueError("INT8 VAE epilogue scales must be on the accumulator device") + if bias is not None and (bias.device != acc.device or bias.shape != (cols,) or not bias.is_contiguous()): + raise ValueError("INT8 VAE epilogue bias must be contiguous on the accumulator device") + if dtype not in (torch.float32, torch.float16, torch.bfloat16): + raise ValueError("INT8 VAE epilogue supports FP32, FP16 and BF16 outputs") + out = torch.empty((rows, cols), device=acc.device, dtype=dtype) + _dequant_bias[(triton.cdiv(rows * cols, 1024),)]( + acc, x_scale, weight_scale, bias if bias is not None else acc, out, + rows, cols, x_scale.stride(0), weight_scale.stride(0), bias is not None, + BLOCK=1024, num_warps=4, enable_fp_fusion=False) + return out diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index e116defb7f..fca78c244a 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -15,6 +15,8 @@ import torch.nn.functional as F from torch.utils.checkpoint import checkpoint +import fastvideo.envs as envs + from fastvideo.attention import get_attn_backend from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig from fastvideo.platforms import AttentionBackendEnum @@ -300,6 +302,7 @@ def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: self.heads = heads self.dim_head = dim_head self.use_bias = bias + self._share_int8_qkv = envs.FASTVIDEO_H3_VAE_INT8_SHARED_QKV.get() inner_dim = heads * dim_head self.norm_q = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=False) self.norm_k = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=False) @@ -339,9 +342,13 @@ def forward( 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)) + if self._share_int8_qkv and not torch.is_grad_enabled() and not torch.compiler.is_compiling(): + from fastvideo.models.vaes.minimax_h3_int8_convrot import shared_int8_projections + projections = shared_int8_projections((self.to_q, self.to_k, self.to_v), hidden_states) + else: + projections = tuple(layer(hidden_states) for layer in (self.to_q, self.to_k, self.to_v)) + query, key, value = (projection.unflatten(2, (self.heads, -1)) for projection in projections) + del projections query = self.norm_q(query.float()).to(query.dtype) key = self.norm_k(key.float()).to(key.dtype) @@ -522,6 +529,11 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: ) +def _tile_batch_size() -> int: + """Spatial tiles decoded per decoder call (``FASTVIDEO_H3_VAE_TILE_BATCH``, default 1 = per tile).""" + return max(1, envs.FASTVIDEO_H3_VAE_TILE_BATCH.get()) + + 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) @@ -817,10 +829,31 @@ def _decode_clip(self, z: torch.Tensor) -> torch.Tensor: ratio = self.spatial_compression_ratio rows = [] + if _tile_batch_size() > 1 and len(set(y_lengths)) == 1 and len(set(x_lengths)) == 1: + # Every tile of the grid has one shape, and the ViT decoder + # treats batch entries independently: decode the grid in a few + # large calls instead of one small call per tile. + with nvtx_range("minimax_h3.vae.decode_clip.decode_tile_batches"): + latent_tiles = [ + z[..., y_position // ratio:(y_position + y_length) // ratio, + x_position // ratio:(x_position + x_length) // ratio] + for y_position, y_length in zip(y_indices, y_lengths) + for x_position, x_length in zip(x_indices, x_lengths) + ] + decoded: list[torch.Tensor] = [] + per_call = _tile_batch_size() + for start in range(0, len(latent_tiles), per_call): + batch = torch.cat(latent_tiles[start:start + per_call], dim=0) + # Clone before retaining: under the reduce-overhead compile the decoder output lives in + # a CUDA-graph pool that the next batch's replay overwrites. + out = self.decoder(self._project_decoder_tile(batch)).clone() + decoded.extend(out.split(z.shape[0], dim=0)) + columns = len(x_indices) + rows = [decoded[index:index + columns] for index in range(0, len(decoded), columns)] # 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)): + for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths) if not rows else ()): 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}"): diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index b39f5c443c..f56a926af9 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -"""FastVideo composed pipelines for MiniMax H3.""" from __future__ import annotations @@ -11,6 +10,9 @@ from typing import Any import torch +from torch.distributed.tensor import DTensor + +import fastvideo.envs as envs from fastvideo.attention.selector import (_active_component_attention_backend_scope, coerce_attn_backend, get_env_variable_attn_backend) @@ -19,6 +21,7 @@ from fastvideo.configs.pipelines.minimax_h3 import (FASTH3_INFERENCE_FILE, FASTH3_INFERENCE_SCHEMA, MiniMaxH3PipelineConfig) from fastvideo.fastvideo_args import FastVideoArgs +from fastvideo.hooks.pinned_memory import PinnedTensorArena from fastvideo.logger import init_logger from fastvideo.models.hf_transformer_utils import get_diffusers_config from fastvideo.pipelines.basic.minimax_h3.stages import ( @@ -29,6 +32,7 @@ MiniMaxH3LatentPreparationStage, MiniMaxH3VideoDecodingStage, ) +from fastvideo.pipelines.basic.minimax_h3.vsa_guard import refuse_zero_initialized_h3_vsa from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase from fastvideo.pipelines.lora_pipeline import LoRAPipeline from fastvideo.pipelines.pipeline_batch_info import ForwardBatch @@ -36,9 +40,6 @@ logger = init_logger(__name__) -# Same split as the MLX runtime: condition, release the ~66 GB Qwen3-VL stack, -# then load DiT + VAEs. Keeping them resident together OOMs unified-memory -# boxes (GB10 / Spark) even though host offload is correctly disabled there. _DENOISE_MODULE_NAMES = ("vae", "audio_vae", "transformer") @@ -77,6 +78,48 @@ def _checkpoint_has_vsa_gates(transformer_dir: Path) -> bool: return False +def _pinned_swap(module: Any, device: torch.device) -> None: + """Move a module's tensors between the GPU and a persistent pinned host copy. + + Inference weights never change, so a parameter's pinned copy is made once and parking just repoints the + parameter at it (no transfer); restoring is one pinned host-to-device copy. Buffers are copied every time. + """ + store = module.__dict__.setdefault("_pinned_host_tensors", {}) + params = dict(module.named_parameters()) + tensors = list(params.items()) + list(module.named_buffers()) + if device.type == "cpu": + missing = [(name, tensor) for name, tensor in tensors + if tensor is not None and tensor.device.type != "cpu" and ( + name not in store or store[name].shape != tensor.shape or store[name].dtype != tensor.dtype)] + arena = PinnedTensorArena(missing) if missing else None + for name, tensor in tensors: + if tensor is None: + continue + if device.type == "cpu": + if tensor.device.type == "cpu": + continue + host = store.get(name) + if host is None or host.shape != tensor.shape or host.dtype != tensor.dtype: + assert arena is not None + host = arena.empty_like(name, tensor) + host.copy_(tensor) + store[name] = host + elif name not in params: + host.copy_(tensor) + tensor.data = host + elif tensor.device != device: + tensor.data = tensor.data.to(device, non_blocking=True) + if device.type != "cpu" and torch.cuda.is_available(): + torch.cuda.current_stream(device).synchronize() + + +def _module_has_dtensor_params(module: Any) -> bool: + parameters = getattr(module, "parameters", None) + if not callable(parameters): + return False + return any(isinstance(parameter, DTensor) for parameter in parameters()) + + @dataclass(frozen=True) class _H3VideoGeometry: spatial_compression_ratio: int @@ -102,7 +145,6 @@ def _default_audio_geometry() -> _H3AudioGeometry: def _apply_h3_checkpoint_arch_configs(model_path: str, fastvideo_args: FastVideoArgs, extra_config_module_map: dict[str, str]) -> None: - """Overlay checkpoint config.json onto pipeline configs without loading weights.""" root = Path(model_path) vae_dir = root / extra_config_module_map.get("vae", "vae") if (vae_dir / "config.json").is_file(): @@ -131,16 +173,7 @@ def _use_taeh3_t2va(fastvideo_args: FastVideoArgs | None, *, ref2va: bool) -> bo class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): - """Shared loading and target-generation path for MiniMax H3. - - Inherits ``LoRAPipeline`` so acceleration and distillation adapters can be merged - in; without it every adapter is rejected with "pipeline is not a LoRAPipeline". - """ - # The linears every published H3 adapter targets. Left unset, ``LoRAPipeline`` - # wraps *every* linear in the DiT -- including ``proj_in``, whose ``.weight`` the - # forward pass reads directly. ``BaseLayerWithLoRA`` exposes no ``.weight``, so - # that wrapping turns generation into an AttributeError before the first step. lora_target_modules = [ "attn.to_q", "attn.to_k", @@ -149,8 +182,6 @@ class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): "ff.fc_in", "ff.fc_out", "adaln_proj.linear", - # The final AdaLN projection. Published community adapters (larryvrh's Turbo) - # target it as `final_layer.adaln_proj.linear`. "norm_out.linear", ] @@ -166,10 +197,6 @@ class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): "scheduler", "audio_scheduler", ] - # Deferral is safe here: geometry scalars come from checkpoint config.json - # (applied in initialize_pipeline without loading weights), no stage - # constructor reads a deferred component, and initialize_pipeline only - # inspects the schedulers, which are never deferred. _lazy_module_names = ("text_encoder", "transformer", "vae", "audio_vae") def __init__(self, *args: Any, **kwargs: Any) -> None: @@ -191,14 +218,14 @@ def get_hf_download_allow_patterns(cls) -> list[str]: def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None: _apply_h3_checkpoint_arch_configs(self.model_path, fastvideo_args, self._extra_config_module_map) - # Each modality's scheduler_config.json owns its shift. Base H3 keeps - # 12/3; a distilled checkpoint can serialize a different trained pair - # (for example 10/3) without being silently rewritten to base defaults. for module_name, modality in (("scheduler", "video"), ("audio_scheduler", "audio")): shift = getattr(self.get_module(module_name), "shift", None) if shift is None or not math.isfinite(float(shift)) or float(shift) <= 0: raise ValueError(f"MiniMax-H3 {modality} scheduler must expose a positive finite shift, got {shift}.") self._load_checkpoint_schedule(fastvideo_args) + transformer = self.get_module("transformer") + if transformer is not None: + refuse_zero_initialized_h3_vsa(transformer) def _checkpoint_facts(self) -> tuple[Any, bool]: """The checkpoint's parsed ``fastvideo_inference.json`` (None without one) and whether its transformer @@ -310,13 +337,6 @@ def _load_checkpoint_schedule(self, fastvideo_args: FastVideoArgs) -> None: def _defer_denoise_modules(self, fastvideo_args: FastVideoArgs) -> bool: if not fastvideo_args.inference_mode or bool(getattr(fastvideo_args, "training_mode", False)): return False - # Both mechanisms defer the same four modules and both decide when to - # free them. Running them together strips DiT/VAEs from the first load - # (sequential) while the base wraps the encoder in a proxy (lazy), so - # post_init's VAE compile transform has nothing to attach to. Lazy is - # the more general owner — including auto-on for unified memory — so it - # wins whenever it is on. Sequential remains the H3-only fallback when - # lazy is off. if bool(getattr(fastvideo_args, "lazy_module_load", False)): logger.info("MiniMax-H3 sequential module load off: lazy_module_load owns deferral") return False @@ -363,8 +383,6 @@ def load_modules(self, return super().load_modules(fastvideo_args, loaded_modules) saved = list(self.required_config_modules) - # Always defer the full denoise set on the first load. TAEH3 T2VA then - # omits the video VAE from the second load via `_denoise_module_names`. self._required_config_modules = [name for name in saved if name not in _DENOISE_MODULE_NAMES] try: logger.info("Loading MiniMax-H3 condition modules first: %s", self._required_config_modules) @@ -390,24 +408,46 @@ def _load_denoise_modules(self, fastvideo_args: FastVideoArgs) -> None: finally: self._required_config_modules = saved + def _unified_memory_host(self) -> bool: + from fastvideo.pipelines import composed_pipeline_base + from fastvideo.platforms import current_platform + + device = composed_pipeline_base.get_local_torch_device() + device_id = 0 if device.index is None else int(device.index) + return bool(current_platform.has_unified_memory(device_id)) + def _release_text_encoder(self) -> None: - stage = self._stage_name_mapping.get("conditioning_stage") - if stage is not None: - stage.conditioner = None - encoder = self.modules.pop("text_encoder", None) + encoder = self.get_module("text_encoder") if encoder is None: return - logger.info("Released MiniMax-H3 text encoder after conditioning") - del encoder + # Unified-memory boxes cannot keep Qwen around even on "CPU". Discrete + # GPUs can: pin it in host RAM and borrow the GPU only for encode. + if self._unified_memory_host(): + stage = self._stage_name_mapping.get("conditioning_stage") + if stage is not None: + stage.conditioner = None + self.modules.pop("text_encoder", None) + logger.info("Released MiniMax-H3 text encoder after conditioning") + del encoder + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return + if not self._move_module(encoder, "cpu"): + return gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() + logger.info("Moved MiniMax-H3 text encoder to CPU after conditioning") def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: - """Reload Qwen3-VL after `_release_text_encoder` so a later request can encode.""" encoder = self.get_module("text_encoder") stage = self._stage_name_mapping.get("conditioning_stage") if encoder is not None: + if not self._unified_memory_host(): + from fastvideo.pipelines import composed_pipeline_base + + self._move_module(encoder, composed_pipeline_base.get_local_torch_device()) if stage is not None and getattr(stage, "conditioner", None) is None: stage.conditioner = encoder return @@ -424,11 +464,70 @@ def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: if stage is not None: stage.conditioner = self.get_module("text_encoder") + def _move_module(self, module: Any, device: str | torch.device) -> bool: + if getattr(module, "_h3_encoder_layerwise_device", None) is not None: + # Layer hooks own placement; moving the whole encoder would restore + # every weight at once and defeat its VRAM bound. + return True + if _module_has_dtensor_params(module): + return False + if envs.FASTVIDEO_H3_PINNED_SWAP.get(): + _pinned_swap(module, torch.device(device)) + else: + module.to(device) + return True + + @staticmethod + def _parked_module_names() -> tuple[str, ...]: + """Denoise modules parked on the host while the text encoder runs (FASTVIDEO_H3_PARK_MODULES). + + Cards with room for the DiT next to the encoder park only the VAEs and keep the DiT resident. + """ + requested = envs.FASTVIDEO_H3_PARK_MODULES.get() + if not requested: + return _DENOISE_MODULE_NAMES + return tuple(name for name in requested.split(",") if name in _DENOISE_MODULE_NAMES) + + def _park_denoise_modules(self) -> None: + parked = False + for name in self._parked_module_names(): + module = self.get_module(name) + if module is None: + continue + if self._move_module(module, "cpu"): + parked = True + if parked: + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + logger.info("Parked MiniMax-H3 denoise modules on CPU for text encode") + + def _restore_denoise_modules(self, fastvideo_args: FastVideoArgs) -> None: + from fastvideo.pipelines import composed_pipeline_base + + device = composed_pipeline_base.get_local_torch_device() + restored = False + for name in _DENOISE_MODULE_NAMES: + # Encode/decode stages move each VAE to the device when consumed. + # Keeping offloaded VAEs on the host leaves room for DiT activations + # and resident blocks throughout the denoising loop. + if name in {"vae", "audio_vae"} and fastvideo_args.vae_cpu_offload: + continue + module = self.get_module(name) + if module is None: + continue + if self._move_module(module, device): + restored = True + if restored: + logger.info("Restored MiniMax-H3 denoise modules to %s", device) + def _run_condition_then_denoise(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: for name in ("input_preparation_stage", "conditioning_stage"): batch = self._stage_name_mapping[name](batch, fastvideo_args) self._release_text_encoder() self._load_denoise_modules(fastvideo_args) + if not self._unified_memory_host(): + self._restore_denoise_modules(fastvideo_args) if not self._denoise_stages_ready: self._add_denoise_stages(ref2va=self._ref2va) for name in ( @@ -441,7 +540,6 @@ def _run_condition_then_denoise(self, batch: ForwardBatch, fastvideo_args: FastV return batch def _input_video_geometry(self, fastvideo_args: FastVideoArgs) -> Any: - """Read canvas scalars from checkpoint JSON, not a live VAE proxy.""" arch = getattr(getattr(fastvideo_args.pipeline_config, "vae_config", None), "arch_config", None) if arch is not None: return arch @@ -482,6 +580,8 @@ def _add_condition_stages(self, fastvideo_args: FastVideoArgs, *, ref2va: bool) def _add_denoise_stages(self, *, ref2va: bool) -> None: transformer = self.get_module("transformer") + if transformer is not None: + refuse_zero_initialized_h3_vsa(transformer) vae = self.get_module("vae") audio_vae = self.get_module("audio_vae") scheduler = self.get_module("scheduler") @@ -523,13 +623,10 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward if not self.post_init_called: self.post_init() - # Sequential encode-then-release is the H3-only fallback. Lazy and the - # fully-resident discrete-GPU path both keep a complete stage list and - # must use the base forward so abort cleanup and text_encoder_cpu_offload - # still apply. Releasing Qwen on every request was re-reading it from disk - # when neither deferral flag was on. if self._defer_denoise_modules(fastvideo_args): try: + if not self._unified_memory_host(): + self._park_denoise_modules() self._ensure_text_encoder(fastvideo_args) if self._denoise_stages_ready: logger.info("Running MiniMax-H3 condition stages before denoise (subsequent request)") diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py index d39128e91d..1e26ee5409 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py @@ -261,7 +261,8 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward 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 - and not isinstance(first_param, DTensor)) + and not isinstance(first_param, DTensor) + and getattr(self.conditioner, "_h3_encoder_layerwise_device", None) is None) if moved_for_forward: self.conditioner.to(device) try: diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py index c59ac6f113..83d9a004a0 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py @@ -22,6 +22,8 @@ build_row_timesteps, ) from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_latent_preparation import MINIMAX_H3_LAYOUT_KEY +from fastvideo.pipelines.basic.minimax_h3.vsa_guard import refuse_zero_initialized_h3_vsa +from fastvideo.pipelines.lazy_module import is_lazy_module from fastvideo.pipelines.pipeline_batch_info import ForwardBatch from fastvideo.pipelines.stages.base import PipelineStage from fastvideo.pipelines.stages.validators import StageValidators as V @@ -191,6 +193,10 @@ def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> V @torch.no_grad() def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: """Denoise the packed H3 video and audio streams over one shared schedule.""" + transformer = self.transformer + if is_lazy_module(transformer): + transformer.materialize() + refuse_zero_initialized_h3_vsa(transformer) layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY) if not isinstance(layout, MiniMaxH3PackedLayout): raise ValueError("MiniMax-H3 packed layout is missing before denoising.") diff --git a/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py new file mode 100644 index 0000000000..5820252acd --- /dev/null +++ b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py @@ -0,0 +1,42 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +from fastvideo.pipelines.lazy_module import is_lazy_module + + +def refuse_zero_initialized_h3_vsa(transformer: Any) -> None: + if is_lazy_module(transformer) and not transformer.is_materialized: + return + blocks = getattr(transformer, "transformer_blocks", None) + if blocks is None: + return + saw_weight = False + any_trained = False + for block in blocks: + attention = getattr(block, "attn", None) + gate = getattr(attention, "to_gate_compress", None) if attention is not None else None + if gate is None: + continue + weight = getattr(gate, "weight", None) + packed = getattr(gate, "_nvfp4_weight", None) + if (weight is None or weight.numel() == 0) and packed is not None and packed.numel() > 0: + # h3_dit_vsa: the gate is packed NVFP4 (weight is None). Two E2M1 codes per byte; the low three + # bits of each nibble are the magnitude, so a gate whose codes are all +-0 is a zero gate. + saw_weight = True + if bool(((packed & 0x77) != 0).any()): + any_trained = True + break + continue + # Layerwise offload leaves a zero-element placeholder on the module; it says nothing about the gate. + if weight is None or weight.numel() == 0: + continue + saw_weight = True + if bool((weight != 0).any()): + any_trained = True + break + if saw_weight and not any_trained: + raise RuntimeError("VIDEO_SPARSE_ATTN_H3 is loaded but every to_gate_compress weight is zero. " + "Load a VSA-trained student or a LoRA that carries gates; dense CompactH3 cannot run VSA.") diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index 838ce677ad..aa086a7d8f 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -188,6 +188,14 @@ def _execute( torch.cuda.synchronize() execution_time = time.perf_counter() - start_time logger.info("[%s] Execution completed in %s ms", stage_name, execution_time * 1000) + if torch.cuda.is_available(): + gib = 1024**3 + logger.info("[%s] Memory peak_allocated=%.2f GiB reserved=%.2f GiB resident_after=%.2f GiB", + stage_name, + torch.cuda.max_memory_allocated() / gib, + torch.cuda.memory_reserved() / gib, + torch.cuda.memory_allocated() / gib) + torch.cuda.reset_peak_memory_stats() batch.logging_info.add_stage_execution_time(stage_key, execution_time) batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name) component_metric = self.performance_component_metric diff --git a/fastvideo/tests/api/test_typed_quant_flow.py b/fastvideo/tests/api/test_typed_quant_flow.py index b7820d987d..1a3c6cf9aa 100644 --- a/fastvideo/tests/api/test_typed_quant_flow.py +++ b/fastvideo/tests/api/test_typed_quant_flow.py @@ -56,6 +56,21 @@ def test_typed_transformer_quant_resolves_to_nvfp4_instance(captured_kwargs) -> assert isinstance(captured_kwargs["transformer_quant"], NVFP4Config), (f"Expected NVFP4Config instance, got " f"{type(captured_kwargs['transformer_quant']).__name__}") + assert captured_kwargs["transformer_quant"].layer_profile == "refine" + + +def test_typed_nvfp4_layer_profile_pins_h3_dit(captured_kwargs) -> None: + from fastvideo.layers.quantization.nvfp4_config import NVFP4Config as CurrentNVFP4Config + + cfg = GeneratorConfig( + model_path="FastVideo/LTX2-Distilled-Diffusers", + engine=EngineConfig(quantization=QuantizationConfig(transformer_quant="NVFP4", layer_profile="h3_dit"), ), + ) + generator_config_to_fastvideo_args(cfg) + quant = captured_kwargs["transformer_quant"] + assert type(quant).__name__ == "NVFP4Config" + assert getattr(quant, "layer_profile") == "h3_dit" + assert isinstance(quant, CurrentNVFP4Config) def test_no_typed_quant_omits_transformer_quant_kwarg(captured_kwargs) -> None: diff --git a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py new file mode 100644 index 0000000000..431c5279d4 --- /dev/null +++ b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py @@ -0,0 +1,113 @@ +# SPDX-License-Identifier: Apache-2.0 +"""sm89 INT8-QK/FP8-PV regression against dense masked BF16 attention.""" +from __future__ import annotations + +import pytest +import torch + + +def _cuda_sm89(): + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (8, 9): + pytest.skip("RTX 4090 / sm89 CUDA is required") + + +@pytest.mark.parametrize("partial", [False, True]) +def test_sparse_int8_preserves_tile_selection_and_valid_keys(partial): + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + + torch.manual_seed(42) + q, k, v = (torch.randn(1, 2, 256, 128, device="cuda", dtype=torch.bfloat16) for _ in range(3)) + vbs = torch.tensor([64, 7 if partial else 64, 31 if partial else 64, 64], device="cuda", dtype=torch.int32) + valid = torch.arange(256, device="cuda") % 64 < vbs.repeat_interleave(64) + k[..., ~valid, :] = 0 + v[..., ~valid, :] = 0 + # Adjacent query tiles deliberately select different keys. A paired-query + # OR adapter would fail this regression even with perfect quantization. + mask = torch.tensor([[1, 0, 0, 1], [0, 1, 0, 0], [1, 0, 1, 0], [0, 0, 1, 1]], + device="cuda", dtype=torch.bool)[None, None].expand(1, 2, -1, -1).contiguous() + dense_mask = mask.repeat_interleave(64, -2).repeat_interleave(64, -1) & valid[None, None, None, :] + with torch.inference_mode(): + expected = torch.nn.functional.scaled_dot_product_attention(q.float(), k.float(), v.float(), + attn_mask=dense_mask) + output = sparse_sm89_attention(q, k, v, mask, vbs) + assert torch.isfinite(output).all() + relative_error = (output.float() - expected).norm() / expected.norm() + assert relative_error < 0.055, float(relative_error) + torch.testing.assert_close(output.float(), expected, atol=0.05, rtol=0.15) + + +def test_sparse_int8_handles_empty_selection(): + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + + q = torch.zeros(1, 1, 128, 128, device="cuda", dtype=torch.bfloat16) + with torch.inference_mode(): + out = sparse_sm89_attention(q, q, q, torch.zeros(1, 1, 2, 2, device="cuda", dtype=torch.bool), + torch.tensor([64, 64], device="cuda", dtype=torch.int32)) + assert torch.count_nonzero(out) == 0 + + +def test_fp8_dynamic_probability_scale_preserves_small_blocks(): + """A large earlier max must not erase a later block's small P but large V.""" + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + + q = torch.zeros(1, 1, 128, 128, device="cuda", dtype=torch.bfloat16) + k, v = torch.zeros_like(q), torch.zeros_like(q) + q[..., 0] = 16 + k[..., :64, 0] = 14 + k[..., 64:, 0] = 2 + v[..., 64:, :] = 1e7 + mask = torch.ones(1, 1, 2, 2, device="cuda", dtype=torch.bool) + vbs = torch.tensor([64, 64], device="cuda", dtype=torch.int32) + with torch.inference_mode(): + reference = torch.nn.functional.scaled_dot_product_attention(q.float(), k.float(), v.float()) + fixed = sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=False, fp8_pv=True, + fp8_v_tiles=True, fp8_dynamic_p=False) + dynamic = sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=False, fp8_pv=True, + fp8_v_tiles=True, fp8_dynamic_p=True) + assert reference.abs().min() > 0.1 + assert torch.count_nonzero(fixed) == 0 + torch.testing.assert_close(dynamic.float(), reference, rtol=0.02, atol=0.02) + + +@pytest.mark.parametrize("batch,heads", [(1, 2), (2, 3)]) +@pytest.mark.parametrize("partner_pad", [False, True]) +def test_int8_bshd_views_match_contiguous_and_reduce_peak(batch, heads, partner_pad): + """Read production BSHD views without retaining three BHSD copies.""" + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + + torch.manual_seed(113) + length, dim = 1024, 128 + storage_length = length + (64 if partner_pad else 0) + tensors = [torch.randn(batch, storage_length, heads, dim, device="cuda", dtype=torch.bfloat16) + for _ in range(3)] + q, k, v = [tensor[:, :length].transpose(1, 2) for tensor in tensors] + vbs = torch.full((length // 64,), 64, device="cuda", dtype=torch.int32) + vbs[1], vbs[4] = 7, 31 + valid = torch.arange(length, device="cuda") % 64 < vbs.repeat_interleave(64) + k[:, :, ~valid] = 0 + v[:, :, ~valid] = 0 + mask = torch.rand(batch, heads, length // 64, length // 64, device="cuda") > 0.8 + mask[:, :, 0] = False + with torch.inference_mode(): + # Populate autotuning/compilation caches before measuring allocations. + warm = sparse_sm89_attention(q, k, v, mask, vbs) + del warm + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + reference = sparse_sm89_attention(q.contiguous(), k.contiguous(), v.contiguous(), mask, vbs) + torch.cuda.synchronize() + copy_peak = torch.cuda.max_memory_allocated() - baseline + expected = reference.cpu() + del reference + torch.cuda.reset_peak_memory_stats() + actual = sparse_sm89_attention(q, k, v, mask, vbs) + torch.cuda.synchronize() + view_peak = torch.cuda.max_memory_allocated() - baseline + torch.testing.assert_close(actual.cpu(), expected, rtol=0, atol=0) + tensor_bytes = batch * heads * length * dim * 2 + assert copy_peak - view_peak >= tensor_bytes, (copy_peak, view_peak) diff --git a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py new file mode 100644 index 0000000000..abd6e386a0 --- /dev/null +++ b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py @@ -0,0 +1,125 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Text-only H3 encoder streaming parity and placement contracts.""" +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch + +import fastvideo.envs as envs + +from fastvideo.hooks.hooks import ModuleHookManager +from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLArchConfig, MiniMaxH3Qwen3VLConfig +from fastvideo.models.encoders.minimax_h3_checkpoint_nvfp4 import MiniMaxH3SerializedNVFP4Config +from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner +from fastvideo.models.loader.text_encoder_quantization import _process_quantized_text_encoder_weights +from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import MiniMaxH3Pipeline +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage +from fastvideo.pipelines.pipeline_batch_info import ForwardBatch + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for encoder streaming") +@pytest.mark.parametrize("quantized,fused", [(False, False), (True, False), (True, True)]) +@pytest.mark.parametrize("pin_cpu_memory,checkpoint_backed", [(True, False), (False, False), (False, True)]) +def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, env_overrides, quantized, + fused, pin_cpu_memory, checkpoint_backed, tmp_path): + # DiT residency must not accidentally keep encoder layers resident too. + env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.override(6)) + env_overrides.enter_context(envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.override(False)) + config = MiniMaxH3Qwen3VLConfig() + config.arch_config = MiniMaxH3Qwen3VLArchConfig( + vocab_size=64, hidden_size=128, intermediate_size=256, + num_hidden_layers=3, num_hidden_layers_override=2, output_hidden_state_index=2, + num_attention_heads=1, num_key_value_heads=1, head_dim=128, + rope_scaling={"mrope_interleaved": True, "mrope_section": [32, 16, 16], "rope_type": "default"}, + vision_depth=1, vision_hidden_size=64, vision_intermediate_size=128, + vision_num_heads=1, vision_deepstack_visual_indexes=(), vision_out_hidden_size=128, + ) + config.quant_config = MiniMaxH3SerializedNVFP4Config() if quantized else None + torch.manual_seed(81) + model = MiniMaxH3Qwen3VLConditioner(config).to(dtype=torch.bfloat16).eval() + for name, parameter in model.named_parameters(): + if name.endswith("weight_packed"): + parameter.data.random_(0, 256) + elif name.endswith("weight_scale"): + parameter.data.fill_(0x38) + elif name.endswith("weight_global_scale"): + parameter.data.fill_(2.7) + else: + parameter.data.normal_(std=0.02) + if quantized: + _process_quantized_text_encoder_weights(model, torch.device("cuda")) + linear = model.language_model.layers[0].self_attn.q_proj.to("cuda") + x = torch.randn(3, 128, device="cuda", dtype=torch.bfloat16) + expected_linear = linear(x)[0] + with patch.object(torch.Tensor, "item", side_effect=AssertionError("Unexpected device scalar read")): + actual_linear = linear(x)[0] + torch.testing.assert_close(actual_linear, expected_linear, rtol=0, atol=0) + checkpoint = tmp_path / "encoder.safetensors" + if checkpoint_backed: + from safetensors.torch import save_file + save_file({name: parameter.detach().cpu().contiguous() for name, parameter in model.named_parameters()}, + checkpoint) + ids = torch.tensor([1, 7, 4, 21, 5, 31, 18], device="cuda") + model.to("cuda") + expected = model.encode_ids(ids) + assert torch.isfinite(expected).all() + if fused: + for layer in model.modules(): + if hasattr(layer, "_nvfp4_fused_dequant"): + layer._nvfp4_fused_dequant = True + model.to("cpu") + mapped = {} + if checkpoint_backed: + from safetensors.torch import load_file + mapped = load_file(checkpoint, device="cpu") + model.enable_checkpoint_backed_cpu_load() + model.load_weights(mapped.items()) + if quantized: + _process_quantized_text_encoder_weights(model, torch.device("cpu")) + for name, parameter in model.named_parameters(): + assert parameter.data_ptr() == mapped[name].data_ptr() + model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) + model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) # repeated setup is harmless + assert model.language_model.embed_tokens.weight.device.type == "cpu" + assert next(model.visual.parameters()).device.type == "cpu" + for _ in range(2): + actual = model.encode_ids(ids) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + for layer in model.language_model.layers: + assert all(parameter.numel() == 0 for parameter in layer.parameters()) + manager = ModuleHookManager.get_from(layer) + assert manager is not None + state = manager.forward_hooks["LayerwiseOffloadHook"].state + assert not state.gpu_named_parameters + assert state.pin_cpu_memory == pin_cpu_memory + assert (state.cpu_arena is not None) == pin_cpu_memory + if checkpoint_backed: + for name, tensor in state.cpu_named_parameters.items(): + layer_name = next(key for key, value in model.named_modules() if value is layer) + assert tensor.data_ptr() == mapped[f"{layer_name}.{name}"].data_ptr() + with pytest.raises(ValueError, match="text-only"): + model.encode_ids(ids, pixel_values=torch.zeros(1, device="cuda"), + image_grid_thw=torch.ones(1, 3, device="cuda", dtype=torch.int64)) + + +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +def test_pipeline_does_not_move_streamed_encoder_whole(device): + module = SimpleNamespace(_h3_encoder_layerwise_device=torch.device("cuda")) + module.to = lambda *_: pytest.fail("Whole encoder move defeats streaming") + assert MiniMaxH3Pipeline._move_module(None, module, device) + + +def test_conditioning_stage_keeps_streamed_encoder_placement(monkeypatch): + import fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning as conditioning + + module = SimpleNamespace(_h3_encoder_layerwise_device=torch.device("cuda")) + module.parameters = lambda: iter([torch.empty(1)]) + module.to = lambda *_: pytest.fail("Conditioning must retain layerwise placement") + stage = MiniMaxH3ConditioningStage.__new__(MiniMaxH3ConditioningStage) + stage.conditioner, stage.ref2va = module, False + stage._encode_fl2va = lambda *_: (torch.zeros(1, 2, 128), torch.zeros(2, dtype=torch.int32)) + monkeypatch.setattr(conditioning, "get_local_torch_device", lambda: torch.device("cpu")) + batch = ForwardBatch(data_type="video", prompt="streaming parity") + output = stage.forward(batch, SimpleNamespace(text_encoder_cpu_offload=True)) + assert output.prompt_embeds[0].shape == (1, 2, 128) diff --git a/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py b/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py index a92bb7eb73..7fa85c1f97 100644 --- a/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py +++ b/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py @@ -260,9 +260,13 @@ def test_runtime_preflight_gates_device_capability_parallelism_and_flashinfer(mo with pytest.raises(RuntimeError, match="requires a CUDA device"): config.validate_runtime(torch.device("cpu")) - monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 9)) - with pytest.raises(RuntimeError, match="sm100 or newer"): + monkeypatch.setattr(h3_nvfp4, "get_tp_world_size", lambda: 1) + # Below sm100 there is no FP4 GEMM: sm80+ de-quantizes each linear to bf16, older GPUs are refused. + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (7, 5)) + with pytest.raises(RuntimeError, match="sm80"): config.validate_runtime(torch.device("cuda")) + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 9)) + config.validate_runtime(torch.device("cuda")) monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (12, 1)) monkeypatch.setattr(h3_nvfp4, "get_tp_world_size", lambda: 2) @@ -388,6 +392,7 @@ def fake_linear(x_fp4, x_scale, weight_packed, weight_scale, alpha): receipt.update(alpha=alpha, weight_packed=weight_packed, weight_scale=weight_scale) return torch.zeros(x_fp4.shape[0], weight_packed.shape[0], dtype=torch.bfloat16) + monkeypatch.setattr(h3_nvfp4, "_fp4_gemm_supported", lambda device: True) monkeypatch.setattr(h3_nvfp4, "_quantize_activation_nvfp4", fake_quantize) monkeypatch.setattr(h3_nvfp4, "_nvfp4_linear", fake_linear) output = MiniMaxH3SerializedNVFP4LinearMethod._apply_finalized(layer, torch.ones(1, 5, 128), layer.bias) @@ -402,6 +407,23 @@ def fake_linear(x_fp4, x_scale, weight_packed, weight_scale, alpha): assert torch.equal(output, torch.full((1, 5, 256), 3.0, dtype=torch.bfloat16)) +def test_apply_without_fp4_gemm_matches_the_dequantized_linear(distributed_setup, monkeypatch) -> None: + """Pre-Blackwell GPUs expand the packed weight to bf16 for the call; the result is x @ W_deq.T + b.""" + config = MiniMaxH3SerializedNVFP4Config.from_config(_checkpoint_quantization_config()) + layer = _language_linear(config, bias=True) + _fill_loaded(layer, global_scale=4.0) + layer.quant_method.process_weights_after_loading(layer) + layer.bias.data.fill_(3.0) + monkeypatch.setattr(h3_nvfp4, "_fp4_gemm_supported", lambda device: False) + x = torch.randn(1, 5, 128, dtype=torch.bfloat16) + output = MiniMaxH3SerializedNVFP4LinearMethod._apply_finalized(layer, x, layer.bias) + + weight = h3_nvfp4.dequantize_serialized_nvfp4(layer.weight_packed, layer.weight_scale, 4.0, torch.bfloat16) + # 0x11 packs two E2M1 codes of 0.5; E4M3 block scale 1; global scale 4 -> every weight is 0.125. + assert torch.equal(weight, torch.full_like(weight, 0.125)) + torch.testing.assert_close(output, torch.nn.functional.linear(x, weight, layer.bias.to(torch.bfloat16))) + + def _tiny_conditioner_config(keep_bf16: tuple[str, ...] = ("mlp.down_proj", )) -> MiniMaxH3Qwen3VLConfig: """One language layer whose linears satisfy the FlashInfer tile geometry, and a small vision tower.""" config = MiniMaxH3Qwen3VLConfig() @@ -432,10 +454,13 @@ def _tiny_conditioner_config(keep_bf16: tuple[str, ...] = ("mlp.down_proj", )) - return config -def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) -> None: +@pytest.mark.parametrize("checkpoint_backed", [False, True]) +def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup, checkpoint_backed) -> None: """The real chain: a conditioner built with the NVFP4 config, checkpoint keys spelled the way the converter writes them, ``load_weights``, the strict missing-tensor check, then the post-load hook.""" model = MiniMaxH3Qwen3VLConditioner(_tiny_conditioner_config()) + if checkpoint_backed: + model.enable_checkpoint_backed_cpu_load() layer = model.language_model.layers[0] assert isinstance(layer.self_attn.q_proj.quant_method, MiniMaxH3SerializedNVFP4LinearMethod) assert isinstance(layer.self_attn.o_proj.quant_method, MiniMaxH3SerializedNVFP4LinearMethod) @@ -460,6 +485,9 @@ def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) loaded = model.load_weights(iter(checkpoint.items())) assert loaded == expected assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 6 + if checkpoint_backed: + for name, parameter in model.named_parameters(): + assert parameter.data_ptr() == checkpoint[f"model.{name}"].data_ptr() assert layer.self_attn.q_proj._nvfp4_alpha.item() == pytest.approx(0.5) assert layer.mlp.up_proj._nvfp4_alpha.item() == pytest.approx(0.5) @@ -482,3 +510,32 @@ def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) stray["model.language_model.layers.0.self_attn.q_proj.weight"] = torch.zeros(128, 128) with pytest.raises(ValueError, match="Unexpected"): model.load_weights(iter(stray.items())) + + +def test_checkpoint_backed_load_keeps_custom_loader_semantics(distributed_setup) -> None: + model = MiniMaxH3Qwen3VLConditioner(_tiny_conditioner_config()) + model.enable_checkpoint_backed_cpu_load() + name = "language_model.layers.0.input_layernorm.weight" + parameter = dict(model.named_parameters())[name] + source = torch.ones_like(parameter) + + def custom_loader(target, value): + target.data.copy_(value + 1) + + parameter.weight_loader = custom_loader + model.load_weights([(name, source)]) + torch.testing.assert_close(parameter, source + 1, rtol=0, atol=0) + assert parameter.data_ptr() != source.data_ptr() + + +def test_checkpoint_backed_embedding_rejects_padded_checkpoint_rows(distributed_setup) -> None: + config = _tiny_conditioner_config() + config.arch_config.vocab_size = 63 + model = MiniMaxH3Qwen3VLConditioner(config) + model.enable_checkpoint_backed_cpu_load() + weight = model.language_model.embed_tokens.weight + assert weight.shape[0] == 64 + # Matching the allocated padded shape must not bypass the loader's check + # that checkpoint rows match the original unpadded vocabulary. + with pytest.raises(AssertionError): + model.load_weights([("language_model.embed_tokens.weight", torch.ones_like(weight))]) diff --git a/fastvideo/tests/hooks/test_layerwise_offload.py b/fastvideo/tests/hooks/test_layerwise_offload.py index f1505029e6..ed95b6a0d9 100644 --- a/fastvideo/tests/hooks/test_layerwise_offload.py +++ b/fastvideo/tests/hooks/test_layerwise_offload.py @@ -44,6 +44,36 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return x +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +def test_pageable_offload_retains_mapped_storage_and_matches_repeated_forwards(tmp_path): + """File-backed weights stay reclaimable; GPU prefetch preserves their values.""" + model = SimpleModelWithModuleList(num_blocks=3, hidden_size=32) + pointers = {} + for block_index, block in enumerate(model.blocks): + for name, parameter in block.named_parameters(): + path = tmp_path / f"{block_index}-{name}.bin" + mapped = torch.from_file(str(path), shared=True, size=parameter.numel(), dtype=parameter.dtype) + mapped.copy_(parameter.detach().reshape(-1)) + parameter.data = mapped.view_as(parameter) + pointers[block_index, name] = parameter.data_ptr() + reference = SimpleModelWithModuleList(num_blocks=3, hidden_size=32).cuda() + reference.load_state_dict(model.state_dict()) + x = torch.randn(2, 9, 32, device="cuda") + with torch.inference_mode(): + expected = reference(x) + enable_layerwise_offload(model, pin_cpu_memory=False) + for block_index, block in enumerate(model.blocks): + manager = ModuleHookManager.get_from(block) + state = manager.forward_hooks["LayerwiseOffloadHook"].state + assert state.cpu_arena is None + for name, host in state.cpu_named_parameters.items(): + assert not host.is_pinned() + assert host.data_ptr() == pointers[block_index, name] + with torch.inference_mode(): + for _ in range(3): + torch.testing.assert_close(model(x), expected, rtol=0, atol=0) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") def test_layerwise_offload_basic(): """Test basic functionality of layerwise offloading.""" diff --git a/fastvideo/tests/hooks/test_pinned_memory.py b/fastvideo/tests/hooks/test_pinned_memory.py new file mode 100644 index 0000000000..9402397a8e --- /dev/null +++ b/fastvideo/tests/hooks/test_pinned_memory.py @@ -0,0 +1,109 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CUDA host-registration lifetime and offload regressions (one GPU required).""" + +import gc +import weakref + +import pytest +import torch + +import fastvideo.envs as envs +from torch import nn + +from fastvideo.hooks.hooks import ModuleHookManager +from fastvideo.hooks.layerwise_offload import LayerwiseOffloadHook, LayerwiseOffloadState +from fastvideo.hooks.pinned_memory import PinnedTensorArena + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA host registration requires one GPU") + + +def test_arena_mixed_dtype_exact_copy_and_lifetime(): + sources = { + "weight": torch.randn(17, 33, device="cuda", dtype=torch.bfloat16), + "scale": torch.randn(17, 1, device="cuda", dtype=torch.float32), + "packed": torch.arange(513, device="cuda").to(torch.uint8), + "scalar": torch.tensor(3.0, device="cuda"), + "empty": torch.empty(0, 4, device="cuda"), + } + arena = PinnedTensorArena(sources.items()) + assert arena.buffer is not None, "This GPU must exercise registration, not fallback" + assert arena.buffer.numel() < sum(t.numel() * t.element_size() for t in sources.values()) + 4096 + 256 * len(sources) + hosts = {} + for name, source in sources.items(): + host = arena.empty_like(name, source) + host.copy_(source) + if host.numel(): + assert host.is_pinned() + assert (host.data_ptr() - arena.buffer.data_ptr()) % 256 == 0 or host.numel() == 0 + torch.testing.assert_close(host.to("cuda", non_blocking=True), source, rtol=0, atol=0) + hosts[name] = host + owner = weakref.ref(arena) + del arena + gc.collect() + assert owner() is not None, "Live typed views must retain the registration" + del hosts, host + gc.collect() + assert owner() is None + + +def test_registration_failure_uses_pinned_allocator(monkeypatch): + class RefusingRuntime: + + def cudaHostRegister(self, *_args): + return 1 + + monkeypatch.setattr(torch.cuda, "cudart", lambda: RefusingRuntime()) + source = torch.arange(27, dtype=torch.float32) + arena = PinnedTensorArena([("weight", source)]) + assert arena.buffer is None + host = arena.empty_like("weight", source) + host.copy_(source) + assert host.is_pinned() + torch.testing.assert_close(host, source, rtol=0, atol=0) + arena.close() + arena.close() + + +def test_offload_mutation_and_prefetched_detach(env_overrides): + env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS.override(True)) + module = nn.Linear(16, 16, device="cuda", dtype=torch.bfloat16) + module.register_buffer("packed", torch.arange(1 << 20, device="cuda").to(torch.uint8)) + expected = {name: tensor.clone() for name, tensor in list(module.named_parameters()) + list(module.named_buffers())} + state = LayerwiseOffloadState(torch.cuda.Stream(), torch.device("cuda")) + hook = LayerwiseOffloadHook(state) + manager = ModuleHookManager.get_from_or_default(module) + manager.append_forward_hook(hook) + old_buffer = state.cpu_arena.buffer + assert old_buffer.is_pinned() + with hook.mutate_params_scope(), torch.no_grad(): + module.weight.add_(1) + module.packed.add_(1) + assert not old_buffer.is_pinned(), "Reinitialization must unregister old storage" + expected["weight"].add_(1) + expected["packed"].add_(1) + state.prefetch_params() + new_buffer = state.cpu_arena.buffer + manager.remove_forward_hook(hook.name()) + assert not new_buffer.is_pinned(), "Detachment must unregister storage" + assert state.cpu_arena is None + assert not state.cpu_named_parameters and not state.gpu_named_parameters + for name, tensor in list(module.named_parameters()) + list(module.named_buffers()): + torch.testing.assert_close(tensor, expected[name], rtol=0, atol=0) + + +def test_h3_swap_reuses_host_storage_and_updates_buffers(): + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import _pinned_swap + + module = nn.Linear(16, 16, device="cuda", dtype=torch.bfloat16) + module.register_buffer("cache", torch.ones(11, device="cuda")) + expected_weight = module.weight.detach().clone() + _pinned_swap(module, torch.device("cpu")) + hosts = module._pinned_host_tensors + pointers = {name: tensor.data_ptr() for name, tensor in hosts.items()} + assert all(tensor.is_pinned() for tensor in hosts.values()) + _pinned_swap(module, torch.device("cuda", torch.cuda.current_device())) + module.cache.add_(3) + _pinned_swap(module, torch.device("cpu")) + assert pointers == {name: tensor.data_ptr() for name, tensor in hosts.items()} + torch.testing.assert_close(module.cache, torch.full((11,), 4.0), rtol=0, atol=0) + torch.testing.assert_close(module.weight, expected_weight.cpu(), rtol=0, atol=0) diff --git a/fastvideo/tests/loader/test_cpu_target_checkpoint.py b/fastvideo/tests/loader/test_cpu_target_checkpoint.py new file mode 100644 index 0000000000..093c6bf7b8 --- /dev/null +++ b/fastvideo/tests/loader/test_cpu_target_checkpoint.py @@ -0,0 +1,49 @@ +"""CPU-targeted inference must never stage its checkpoint on the GPU.""" +from types import SimpleNamespace + +import torch +from torch import nn +from safetensors.torch import save_file + +from fastvideo.models.loader import fsdp_load + + +class TinyCPUModel(nn.Module): + param_names_mapping = {} + + def __init__(self, config): + super().__init__() + self.config = config + self.linear = nn.Linear(8, 8) + + +def test_cpu_target_reads_cpu_checkpoint_and_preserves_weights(tmp_path, monkeypatch): + weights = {"linear.weight": torch.arange(64, dtype=torch.float32).view(8, 8), + "linear.bias": torch.arange(8, dtype=torch.float32)} + path = tmp_path / "model.safetensors" + save_file(weights, path) + real_iterator = fsdp_load.safetensors_weights_iterator + placements = [] + + def cpu_iterator(files, *, to_cpu): + assert to_cpu, "CPU target must not stage checkpoint tensors on the GPU" + for name, tensor in real_iterator(files, to_cpu=to_cpu): + placements.append(tensor.device.type) + yield name, tensor + + monkeypatch.setattr(fsdp_load, "safetensors_weights_iterator", cpu_iterator) + model = fsdp_load.maybe_load_fsdp_model( + model_cls=TinyCPUModel, + init_params={"config": SimpleNamespace(quant_config=None)}, + weight_dir_list=[str(path)], device=torch.device("cpu"), + hsdp_replicate_dim=1, hsdp_shard_dim=1, + default_dtype=torch.float32, param_dtype=torch.float32, + reduce_dtype=torch.float32, training_mode=False, cpu_offload=False, + ) + assert placements == ["cpu", "cpu"] + for name, value in model.state_dict().items(): + assert value.device.type == "cpu" + torch.testing.assert_close(value, weights[name], rtol=0, atol=0) + x = torch.randn(2, 8) + torch.testing.assert_close(model.linear(x), torch.nn.functional.linear(x, weights["linear.weight"], + weights["linear.bias"]), rtol=0, atol=0) diff --git a/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py new file mode 100644 index 0000000000..7d34e346f4 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU regression for the ModelOpt activation-scale export contract.""" +import importlib.util +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + + +@pytest.fixture +def converter(monkeypatch): + path = Path(__file__).resolve().parents[4] / 'scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py' + spec = importlib.util.spec_from_file_location('h3_modelopt_converter', path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + packed = torch.zeros((2, 8), dtype=torch.uint8) + scales = torch.ones((2, 1), dtype=torch.float8_e4m3fn) + def quantize(*args, **kwargs): + return packed, scales + monkeypatch.setattr(module, '_flashinfer', lambda: ( + SimpleNamespace(layout_128x4=0), None, quantize, lambda value: value)) + return module + + +def convert(converter, input_scale): + return converter.convert_modelopt_linear( + torch.zeros((2, 8), dtype=torch.uint8), + torch.ones((2, 1), dtype=torch.float8_e4m3fn), + torch.tensor(0.25), 'cpu', input_scale=input_scale)[0] + + +def test_modelopt_activation_scale_is_preserved_as_reciprocal(converter): + buffers = convert(converter, torch.tensor(2.0)) + assert buffers['_nvfp4_input_global_sf'].dtype == torch.float32 + assert buffers['_nvfp4_input_global_sf'].shape == () + assert buffers['_nvfp4_input_global_sf'].item() == 0.5 + assert buffers['_nvfp4_alpha'].item() == 0.25 + + +def test_missing_modelopt_activation_scale_keeps_legacy_export(converter): + assert '_nvfp4_input_global_sf' not in convert(converter, None) + + +@pytest.mark.parametrize('value', [0.0, -1.0, float('nan'), float('inf')]) +def test_invalid_modelopt_activation_scale_is_rejected(converter, value): + with pytest.raises(ValueError, match='finite positive scalar'): + convert(converter, torch.tensor(value)) diff --git a/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py new file mode 100644 index 0000000000..466731b117 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Converter smoke test: a synthetic ModelOpt NVFP4 checkpoint, converted and loaded, matches its reference.""" + +from __future__ import annotations + +import importlib.util +import json +import subprocess +import sys +from pathlib import Path + +import pytest +import torch +import torch.nn as nn +from safetensors import safe_open +from safetensors.torch import save_file + +import fastvideo.layers.quantization.nvfp4_config as nv + +SCRIPT = Path(__file__).resolve().parents[4] / "scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py" +OUT, IN = 128, 128 +ACT_AMAX = 6.0 + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability() < (10, 0), + reason="NVFP4 mm_fp4 needs a Blackwell GPU") + + +def _converter(): + spec = importlib.util.spec_from_file_location("convert_minimax_h3_modelopt_nvfp4_dit", SCRIPT) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _modelopt_linear(prefix: str, gen: torch.Generator) -> dict[str, torch.Tensor]: + """Random E2M1 codes with positive E4M3 block scales: a valid ModelOpt unified-layout linear.""" + return { + f"{prefix}.weight": torch.randint(0, 256, (OUT, IN // 2), dtype=torch.uint8, generator=gen), + f"{prefix}.weight_scale": (torch.rand(OUT, IN // 16, generator=gen) + 0.5).to(torch.float8_e4m3fn), + f"{prefix}.weight_scale_2": torch.tensor(0.01, dtype=torch.float32), + f"{prefix}.input_scale": torch.tensor(1.0, dtype=torch.float32), + } + + +def _write_source(src: Path) -> dict[str, torch.Tensor]: + gen = torch.Generator().manual_seed(0) + tensors = { + **_modelopt_linear("transformer_blocks.0.ff.net.0.proj", gen), + # Already quantized by ModelOpt: --quantize-attention must keep it on the ModelOpt path. + **_modelopt_linear("transformer_blocks.0.attn.to_k", gen), + "transformer_blocks.0.attn.to_q.weight": (torch.randn(OUT, IN, generator=gen) * 0.05).to(torch.bfloat16), + "transformer_blocks.0.norm.weight": torch.ones(IN, dtype=torch.bfloat16), + } + src.mkdir() + save_file(tensors, str(src / "diffusion_pytorch_model-00001-of-00001.safetensors")) + weight_map = {k: "diffusion_pytorch_model-00001-of-00001.safetensors" for k in tensors} + (src / "diffusion_pytorch_model.safetensors.index.json").write_text(json.dumps({"weight_map": weight_map})) + (src / "config.json").write_text(json.dumps({"quantization_config": {"quant_algo": "NVFP4"}})) + return tensors + + +class _Linear(nn.Module): + + def __init__(self, prefix: str) -> None: + super().__init__() + self.weight = nn.Parameter(torch.empty(OUT, IN, dtype=torch.bfloat16), requires_grad=False) + self.quant_method = nv.NVFP4QuantizeMethod(layer_prefix=prefix) + + +def _model() -> nn.Module: + root = nn.Module() + block = nn.Module() + block.ff, block.attn = nn.Module(), nn.Module() + block.ff.fc_in = _Linear("transformer_blocks.0.ff.fc_in") + block.attn.to_q = _Linear("transformer_blocks.0.attn.to_q") + block.attn.to_k = _Linear("transformer_blocks.0.attn.to_k") + root.transformer_blocks = nn.ModuleList([block]) + return root + + +def test_converted_checkpoint_loads_and_matches_reference(tmp_path) -> None: + pytest.importorskip("flashinfer") + conv = _converter() + src, dst = tmp_path / "src", tmp_path / "dst" + source = _write_source(src) + amax = tmp_path / "amax.json" + amax.write_text(json.dumps({f"b0.{name}": {"all": ACT_AMAX} for name in ("ff.fc_in", "attn.to_q", "attn.to_k")})) + subprocess.run([sys.executable, str(SCRIPT), "--src", str(src), "--dst", str(dst), "--quantize-attention", + "--act-amax", str(amax)], check=True) + + export = dst / nv.H3_NVFP4_DIT_EXPORT_FILENAME + with safe_open(str(export), framework="pt", device="cpu") as reader: + keys = set(reader.keys()) + # ModelOpt bytes are carried over bit for bit, including the attention projection it already quantized. + for prefix, module in (("transformer_blocks.0.ff.net.0.proj", "transformer_blocks.0.ff.fc_in"), + ("transformer_blocks.0.attn.to_k", "transformer_blocks.0.attn.to_k")): + assert torch.equal(reader.get_tensor(f"{module}::_nvfp4_weight"), source[f"{prefix}.weight"]) + assert "transformer_blocks.0.attn.to_q::_nvfp4_weight" in keys + dense_index = json.loads((dst / "diffusion_pytorch_model.safetensors.index.json").read_text())["weight_map"] + assert set(dense_index) == {"transformer_blocks.0.norm.weight"} + assert "quantization_config" not in json.loads((dst / "config.json").read_text()) + + model = _model() + assert nv.load_minimax_h3_nvfp4_dit_export(model, str(export), device="cuda") == 3 + x = torch.randn(256, IN, device="cuda").clamp(-ACT_AMAX, ACT_AMAX).to(torch.bfloat16) + references = { + "ff.fc_in": conv.dequantize_modelopt(*(source[f"transformer_blocks.0.ff.net.0.proj.{s}"].cuda() + for s in ("weight", "weight_scale", "weight_scale_2"))), + "attn.to_k": conv.dequantize_modelopt(*(source[f"transformer_blocks.0.attn.to_k.{s}"].cuda() + for s in ("weight", "weight_scale", "weight_scale_2"))), + "attn.to_q": source["transformer_blocks.0.attn.to_q.weight"].cuda(), + } + block = model.transformer_blocks[0] + for name, reference in references.items(): + layer = block.get_submodule(name) + assert layer._nvfp4_input_global_sf.item() == pytest.approx(448.0 * 6.0 / ACT_AMAX) + out = layer.quant_method.apply(layer, x).float() + ref = x.float() @ reference.float().T + error = ((out - ref).norm() / ref.norm()).item() + assert error < 0.2, f"{name}: relative error {error:.3f}" diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_config.py b/fastvideo/tests/ops/quantization/test_nvfp4_config.py index 977528189c..3a9fa0d4dc 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_config.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_config.py @@ -44,10 +44,18 @@ def test_nvfp4config_layer_profile_round_trips_from_dict(): from fastvideo.layers.quantization.nvfp4_config import NVFP4Config config = NVFP4Config.from_config({"layer_profile": "base"}) assert config.layer_profile == "base" + config = NVFP4Config.from_config({"layer_profile": "h3_dit"}) + assert config.layer_profile == "h3_dit" config = NVFP4Config.from_config({}) assert config.layer_profile == "refine" +def test_nvfp4config_rejects_unknown_layer_profile(): + from fastvideo.layers.quantization.nvfp4_config import NVFP4Config + with pytest.raises(ValueError, match="h3_dit"): + NVFP4Config(layer_profile="h3-dit") + + def test_nvfp4_kernel_call_raises_clear_error_without_flashinfer(monkeypatch): """A call into the NVFP4 kernels must raise an actionable ImportError when flashinfer is missing, not a confusing diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_dequant.py b/fastvideo/tests/ops/quantization/test_nvfp4_dequant.py new file mode 100644 index 0000000000..6eb9f6aa87 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_dequant.py @@ -0,0 +1,20 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused weight expansion against the independent serialized Torch decoder.""" +import pytest +import torch + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for Triton NVFP4 decoder") +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +@pytest.mark.parametrize("global_scale", [1.0, 2.7, 438.912]) +def test_fused_nvfp4_expands_all_codes_and_scale_tiles(dtype, global_scale): + from fastvideo.layers.quantization.nvfp4_dequant import dequantize_nvfp4_cuda + from fastvideo.models.encoders.minimax_h3_checkpoint_nvfp4 import dequantize_serialized_nvfp4 + + torch.manual_seed(23) + # Multiple row/column tiles distinguish the swizzle from a row-major decoder. + packed = torch.arange(256, device="cuda", dtype=torch.uint8).repeat(256, 1) + scales = torch.randint(0, 127, (256, 32), device="cuda", dtype=torch.uint8) + reference = dequantize_serialized_nvfp4(packed, scales, global_scale, dtype) + actual = dequantize_nvfp4_cuda(packed, scales, global_scale, dtype) + torch.testing.assert_close(actual, reference, rtol=0, atol=0) diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py b/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py new file mode 100644 index 0000000000..32b9f84d14 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py @@ -0,0 +1,87 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import pytest +import torch +import torch.nn as nn +from safetensors.torch import save_file + +import fastvideo.layers.quantization.nvfp4_config as nv + + +def _method(prefix: str) -> nv.NVFP4QuantizeMethod: + method = object.__new__(nv.NVFP4QuantizeMethod) + method.weight_fp4 = None + method.weight_scale = None + method.x_global_sf = torch.tensor(1.0, dtype=torch.float32) + method.layer_prefix = prefix + method._is_refine_only_layer = False + method._retain_original_weights = None + return method + + +class _ExportLinear(nn.Module): + + def __init__(self, prefix: str, out_dim: int = 8, in_dim: int = 16) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(out_dim, in_dim, dtype=torch.bfloat16), requires_grad=False) + self.quant_method = _method(prefix) + + +def _model() -> nn.Module: + root = nn.Module() + root.transformer_blocks = nn.ModuleList([nn.Module()]) + root.transformer_blocks[0].attn = nn.Module() + root.transformer_blocks[0].attn.to_q = _ExportLinear("transformer_blocks.0.attn.to_q") + return root + + +def test_dense_transformer_safetensors_drops_packed_export(tmp_path) -> None: + packed = tmp_path / nv.H3_NVFP4_DIT_EXPORT_FILENAME + shard = tmp_path / "diffusion_pytorch_model-00001-of-00008.safetensors" + packed.write_bytes(b"") + shard.write_bytes(b"") + kept = nv.dense_transformer_safetensors([str(packed), str(shard)]) + assert kept == [str(shard)] + assert nv.find_minimax_h3_nvfp4_dit_export([str(shard)]) == str(packed) + + +def test_load_minimax_h3_nvfp4_dit_export_overlays_and_purges_weight(tmp_path) -> None: + model = _model() + linear = model.transformer_blocks[0].attn.to_q + prefix = "transformer_blocks.0.attn.to_q" + tensors = { + f"{prefix}::_nvfp4_weight": torch.zeros(8, 8, dtype=torch.uint8), + f"{prefix}::_nvfp4_weight_scale": torch.zeros(8, 1, dtype=torch.uint8), + f"{prefix}::_nvfp4_alpha": torch.tensor(0.5, dtype=torch.float32), + f"{prefix}::_weight_global_sf": torch.tensor(2.0, dtype=torch.bfloat16), + } + path = tmp_path / nv.H3_NVFP4_DIT_EXPORT_FILENAME + save_file(tensors, str(path)) + + loaded = nv.load_minimax_h3_nvfp4_dit_export(model, str(path), device="cpu") + + assert loaded == 1 + assert linear.weight is None + assert linear._nvfp4_weight.dtype is torch.uint8 + assert tuple(linear._nvfp4_weight.shape) == (8, 8) + assert linear._nvfp4_alpha.item() == 0.5 + assert linear._weight_global_sf.item() == 2.0 + + +def test_load_minimax_h3_nvfp4_dit_export_rejects_untagged_linear(tmp_path) -> None: + model = _model() + model.transformer_blocks[0].attn.to_q.quant_method = object() + prefix = "transformer_blocks.0.attn.to_q" + tensors = { + f"{prefix}::_nvfp4_weight": torch.zeros(8, 8, dtype=torch.uint8), + f"{prefix}::_nvfp4_weight_scale": torch.zeros(8, 1, dtype=torch.uint8), + f"{prefix}::_nvfp4_alpha": torch.tensor(1.0, dtype=torch.float32), + f"{prefix}::_weight_global_sf": torch.tensor(1.0, dtype=torch.bfloat16), + } + path = tmp_path / nv.H3_NVFP4_DIT_EXPORT_FILENAME + save_file(tensors, str(path)) + + with pytest.raises(RuntimeError, match="layer_profile='h3_dit'"): + nv.load_minimax_h3_nvfp4_dit_export(model, str(path), device="cpu") diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py b/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py index d5529c04b5..0e03d2daff 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py @@ -47,6 +47,35 @@ def test_is_minimax_h3_nvfp4_linear_prefix_non_main_ffn_linear(prefix: str) -> N assert not nvfp4.is_minimax_h3_nvfp4_linear_prefix(prefix) +@pytest.mark.parametrize( + "prefix", + [ + "transformer_blocks.0.attn.to_q", + "transformer_blocks.41.attn.to_out", + "minimax_h3.transformer_blocks.12.ff.fc_in", + "transformer_blocks.0.attn.to_k", + "transformer_blocks.0.attn.to_v", + "transformer_blocks.3.ff.fc_out", + ], +) +def test_is_minimax_h3_nvfp4_dit_linear_prefix_attn_and_ffn(prefix: str) -> None: + assert nvfp4.is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + + +@pytest.mark.parametrize( + "prefix", + [ + "transformer_blocks.0.adaln_proj.linear", + "token_refiner.refiner_blocks.0.ff.fc_in", + "transformer_blocks.0.attn.to_gate_compress", + "proj_in", + "transformer_blocks.0.attn.to_q.weight", + ], +) +def test_is_minimax_h3_nvfp4_dit_linear_prefix_rejects_non_export_linears(prefix: str) -> None: + assert not nvfp4.is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + + def test_nvfp4config_get_quant_method_minimax_h3_feed_forward() -> None: if not torch.cuda.is_available(): pytest.skip("NVFP4QuantizeMethod construction requires CUDA") diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_purge.py b/fastvideo/tests/ops/quantization/test_nvfp4_purge.py index 385cc47566..18b8566349 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_purge.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_purge.py @@ -114,6 +114,27 @@ def fake_mm_fp4(x_fp4, w_t, x_scale, w_scale_t, alpha, out_dtype, out, backend): assert out.shape == (2, 3, 8) +def test_convert_refuses_fsdp_sharded_bf16_purge(monkeypatch) -> None: + import torch.distributed.tensor as tdt + + class DummyDTensor: + def __init__(self, data: torch.Tensor) -> None: + self._data = data + + def to_local(self) -> torch.Tensor: + return self._data + + def float(self) -> torch.Tensor: + return self._data.float() + + monkeypatch.setattr(tdt, "DTensor", DummyDTensor) + model = _model(retain=False) + del model.always_fp4._parameters["weight"] + object.__setattr__(model.always_fp4, "weight", DummyDTensor(torch.randn(8, 16, dtype=torch.bfloat16))) + with pytest.raises(RuntimeError, match="FSDP-sharded"): + nv.convert_model_to_nvfp4(model) + + def test_dense_path_after_purge_raises_with_flag_named(monkeypatch) -> None: """Defensive guard: convert never purges dense-capable layers anymore, but a hand-purged module hitting the dense path must fail loudly.""" diff --git a/fastvideo/tests/ssim/.reference_videos_download.lock b/fastvideo/tests/ssim/.reference_videos_download.lock new file mode 100644 index 0000000000..e69de29bb2 diff --git a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index 02955b73fe..44138298a8 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -5,8 +5,11 @@ from contextlib import nullcontext from types import SimpleNamespace +import pytest import torch +import fastvideo.envs as envs + import fastvideo.pipelines.composed_pipeline_base as composed_pipeline_base from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.utils import FlexibleArgumentParser @@ -27,19 +30,36 @@ def region(self, name): def _stub_module(name: str) -> SimpleNamespace: if name in {"scheduler"}: - return SimpleNamespace(shift=12.0, name=name) - if name in {"audio_scheduler"}: - return SimpleNamespace(shift=3.0, name=name) - if name == "transformer": + module = SimpleNamespace(shift=12.0, name=name) + elif name in {"audio_scheduler"}: + module = SimpleNamespace(shift=3.0, name=name) + elif name == "transformer": # LoRAPipeline reads exclude_lora_layers off the DiT arch config. - return SimpleNamespace( + module = SimpleNamespace( name=name, config=SimpleNamespace(arch_config=SimpleNamespace(exclude_lora_layers=[])), ) - return SimpleNamespace(name=name) + else: + module = SimpleNamespace(name=name) + + def to(device): + module.device = device + module.moved_to.append(device) + return module + + module.device = None + module.moved_to = [] + module.to = to + return module + + +@pytest.fixture(autouse=True) +def _no_pinned_swap(env_overrides): + # These contract tests use lightweight objects, not tensor-bearing modules. + env_overrides.enter_context(envs.FASTVIDEO_H3_PINNED_SWAP.override(False)) -def _patch_pipeline_construction(monkeypatch, events: list) -> None: +def _patch_pipeline_construction(monkeypatch, events: list, *, unified_memory: bool = False) -> None: monkeypatch.setattr( composed_pipeline_base, "maybe_init_distributed_environment_and_model_parallel", @@ -49,7 +69,7 @@ def _patch_pipeline_construction(monkeypatch, events: list) -> None: monkeypatch.setattr(composed_pipeline_base, "get_world_group", lambda: SimpleNamespace(local_rank=0)) monkeypatch.setattr(composed_pipeline_base, "get_or_create_profiler", lambda trace_dir: _Profiler()) monkeypatch.setattr(composed_pipeline_base, "warmup_sequence_parallel_communication", lambda: None) - monkeypatch.setattr("fastvideo.platforms.current_platform.has_unified_memory", lambda device_id: False) + monkeypatch.setattr("fastvideo.platforms.current_platform.has_unified_memory", lambda device_id: unified_memory) monkeypatch.setattr("fastvideo.platforms.current_platform.is_mps", lambda: False) @@ -103,6 +123,7 @@ def fake_add_denoise(*, ref2va: bool) -> None: monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + encoder = pipeline.get_module("text_encoder") batch = ForwardBatch(data_type="video", prompt="alpine dancer") out = pipeline.forward(batch, args) @@ -111,13 +132,201 @@ def fake_add_denoise(*, ref2va: bool) -> None: assert "transformer" in loads[1] assert "vae" in loads[1] assert "text_encoder" not in loads[1] - assert pipeline.get_module("text_encoder") is None - assert condition_stage.conditioner is None + assert pipeline.get_module("text_encoder") is encoder + assert encoder.moved_to[-1] == "cpu" + assert condition_stage.conditioner is not None assert pipeline.get_module("transformer") is not None assert pipeline._denoise_stages_ready is True + moves_before_second = len(encoder.moved_to) second = pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) assert second is not None + assert len(loads) == 2 + assert pipeline.get_module("text_encoder") is encoder + assert condition_stage.conditioner is not None + assert len(encoder.moved_to) > moves_before_second + assert encoder.moved_to[moves_before_second] == torch.device("cpu") + + +def test_sequential_skips_host_offload_for_dtensor_params(monkeypatch) -> None: + class _FakeDTensor: + pass + + monkeypatch.setattr( + "fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.DTensor", + _FakeDTensor, + ) + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + loads: list[list[str]] = [] + + def _dtensor_stub(name: str) -> SimpleNamespace: + module = _stub_module(name) + + def parameters(): + yield _FakeDTensor() + + module.parameters = parameters + return module + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + if name in modules: + continue + modules[name] = _dtensor_stub(name) if name == "text_encoder" else _stub_module(name) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=True, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + + passthrough = lambda batch, _args: batch + monkeypatch.setattr(pipeline._stage_name_mapping["input_preparation_stage"], "forward", passthrough) + monkeypatch.setattr(pipeline._stage_name_mapping["conditioning_stage"], "forward", passthrough) + original_add_denoise = pipeline._add_denoise_stages + + def fake_add_denoise(*, ref2va: bool) -> None: + original_add_denoise(ref2va=ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + monkeypatch.setattr(pipeline._stage_name_mapping[name], "forward", passthrough) + + monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + encoder = pipeline.get_module("text_encoder") + pipeline.forward(ForwardBatch(data_type="video", prompt="alpine dancer"), args) + pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) + + assert len(loads) == 2 + assert encoder.moved_to == [] + transformer = pipeline.get_module("transformer") + assert transformer is not None + assert transformer.moved_to[-1] == torch.device("cpu") + + +def test_sequential_skips_host_offload_when_dense_params_precede_dtensors(monkeypatch) -> None: + class _FakeDTensor: + pass + + monkeypatch.setattr( + "fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.DTensor", + _FakeDTensor, + ) + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + loads: list[list[str]] = [] + + def _mixed_stub(name: str) -> SimpleNamespace: + module = _stub_module(name) + + def parameters(): + yield torch.zeros(1) + yield _FakeDTensor() + + module.parameters = parameters + return module + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + if name in modules: + continue + modules[name] = _mixed_stub(name) if name == "text_encoder" else _stub_module(name) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=True, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + + passthrough = lambda batch, _args: batch + monkeypatch.setattr(pipeline._stage_name_mapping["input_preparation_stage"], "forward", passthrough) + monkeypatch.setattr(pipeline._stage_name_mapping["conditioning_stage"], "forward", passthrough) + original_add_denoise = pipeline._add_denoise_stages + + def fake_add_denoise(*, ref2va: bool) -> None: + original_add_denoise(ref2va=ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + monkeypatch.setattr(pipeline._stage_name_mapping[name], "forward", passthrough) + + monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + encoder = pipeline.get_module("text_encoder") + pipeline.forward(ForwardBatch(data_type="video", prompt="alpine dancer"), args) + assert encoder.moved_to == [] + + +def test_unified_memory_sequential_deletes_encoder_and_reloads(monkeypatch) -> None: + events: list = [] + _patch_pipeline_construction(monkeypatch, events, unified_memory=True) + loads: list[list[str]] = [] + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + modules.setdefault(name, _stub_module(name)) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=True, + lazy_module_load=False, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + + condition_stage = pipeline._stage_name_mapping["conditioning_stage"] + passthrough = lambda batch, _args: batch + monkeypatch.setattr(pipeline._stage_name_mapping["input_preparation_stage"], "forward", passthrough) + monkeypatch.setattr(condition_stage, "forward", passthrough) + original_add_denoise = pipeline._add_denoise_stages + + def fake_add_denoise(*, ref2va: bool) -> None: + original_add_denoise(ref2va=ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + monkeypatch.setattr(pipeline._stage_name_mapping[name], "forward", passthrough) + + monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + pipeline.forward(ForwardBatch(data_type="video", prompt="alpine dancer"), args) + + assert pipeline.get_module("text_encoder") is None + assert condition_stage.conditioner is None + assert len(loads) == 2 + + pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) assert len(loads) == 3 assert loads[2] == ["text_encoder"] assert pipeline.get_module("text_encoder") is None @@ -323,3 +532,15 @@ def fake_load(self, fastvideo_args, loaded_modules=None): assert first is not None and second is not None assert len(loads) == 1 assert pipeline.get_module("text_encoder") is not None + + +@pytest.mark.parametrize("vae_offload", [True, False]) +def test_sequential_restore_keeps_offloaded_vaes_on_host_until_consumed(monkeypatch, vae_offload): + """Do not occupy denoise VRAM with decoders that stages load on demand.""" + pipeline = MiniMaxH3Pipeline.__new__(MiniMaxH3Pipeline) + pipeline.modules = {name: _stub_module(name) for name in _DENOISE_MODULE_NAMES} + moved = [] + monkeypatch.setattr(composed_pipeline_base, "get_local_torch_device", lambda: torch.device("cuda", 0)) + monkeypatch.setattr(pipeline, "_move_module", lambda module, device: moved.append(module.name) or True) + pipeline._restore_denoise_modules(SimpleNamespace(vae_cpu_offload=vae_offload)) + assert moved == (["transformer"] if vae_offload else list(_DENOISE_MODULE_NAMES)) diff --git a/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py new file mode 100644 index 0000000000..655909f340 --- /dev/null +++ b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py @@ -0,0 +1,88 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch +import torch.nn as nn + +from fastvideo.pipelines.basic.minimax_h3.vsa_guard import refuse_zero_initialized_h3_vsa +from fastvideo.pipelines.lazy_module import LazyModule + + +class _GateAttention(nn.Module): + + def __init__(self, trained: bool, *, has_weight: bool = True) -> None: + super().__init__() + if has_weight: + self.to_gate_compress = nn.Linear(4, 4, bias=False) + with torch.no_grad(): + if trained: + self.to_gate_compress.weight.fill_(0.1) + else: + self.to_gate_compress.weight.zero_() + else: + self.to_gate_compress = SimpleNamespace(weight=None) + + +class _Block(nn.Module): + + def __init__(self, trained: bool, *, has_weight: bool = True) -> None: + super().__init__() + self.attn = _GateAttention(trained, has_weight=has_weight) + + +class _Transformer(nn.Module): + + def __init__(self, trained: bool, *, empty: bool = False, has_weight: bool = True) -> None: + super().__init__() + blocks = [] if empty else [_Block(trained, has_weight=has_weight)] + self.transformer_blocks = nn.ModuleList(blocks) + + +def test_refuse_zero_initialized_h3_vsa_raises() -> None: + with pytest.raises(RuntimeError, match="to_gate_compress"): + refuse_zero_initialized_h3_vsa(_Transformer(trained=False)) + + +def test_refuse_zero_initialized_h3_vsa_allows_trained_gates() -> None: + refuse_zero_initialized_h3_vsa(_Transformer(trained=True)) + + +def test_refuse_zero_initialized_h3_vsa_skips_dense_transformer() -> None: + refuse_zero_initialized_h3_vsa(SimpleNamespace()) + + +def test_refuse_zero_initialized_h3_vsa_skips_empty_blocks() -> None: + refuse_zero_initialized_h3_vsa(_Transformer(trained=False, empty=True)) + + +def test_refuse_zero_initialized_h3_vsa_skips_missing_gate_weight() -> None: + refuse_zero_initialized_h3_vsa(_Transformer(trained=False, has_weight=False)) + + +def test_refuse_zero_initialized_h3_vsa_skips_unmaterialized_lazy_module() -> None: + loads: list[int] = [] + + def loader() -> nn.Module: + loads.append(1) + return _Transformer(trained=False) + + refuse_zero_initialized_h3_vsa(LazyModule("transformer", loader)) + assert loads == [] + + +def _packed_gate_transformer(codes: int) -> SimpleNamespace: + """An h3_dit_vsa gate: ``weight`` purged, two E2M1 codes per ``_nvfp4_weight`` byte.""" + gate = SimpleNamespace(weight=None, _nvfp4_weight=torch.full((4, 2), codes, dtype=torch.uint8)) + return SimpleNamespace(transformer_blocks=[SimpleNamespace(attn=SimpleNamespace(to_gate_compress=gate))]) + + +def test_refuse_zero_initialized_h3_vsa_checks_packed_nvfp4_gates() -> None: + # 0x88 is -0 in both nibbles: still a zero gate. + for zero_codes in (0x00, 0x88): + with pytest.raises(RuntimeError, match="to_gate_compress"): + refuse_zero_initialized_h3_vsa(_packed_gate_transformer(zero_codes)) + refuse_zero_initialized_h3_vsa(_packed_gate_transformer(0x12)) diff --git a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py new file mode 100644 index 0000000000..fa2ac25fcd --- /dev/null +++ b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py @@ -0,0 +1,108 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Tile-first VSA parity on CUDA, including partial tiles and learned gates.""" +from __future__ import annotations + +from unittest.mock import patch + +import pytest +import torch + +import fastvideo.envs as envs + +from fastvideo.layers.quantization.fp8_config import FP8Config, FP8QuantizeMethod +from fastvideo.platforms import AttentionBackendEnum +from fastvideo.models.dits.minimax_h3_vsa_fp4 import _shared_input_projections + + +def _install_fp8_buffers(layer): + weight = layer.weight.data.float() + if layer.quant_method.granularity == "channel": + scale = (weight.abs().amax(dim=1, keepdim=True) / 448).clamp_min(1e-6) + else: + scale = (weight.abs().amax().reshape(1) / 448).clamp_min(1e-6) + layer.register_buffer("_fp8_weight", (weight / scale).to(torch.float8_e4m3fn)) + layer.register_buffer("_fp8_weight_scale", scale) + layer.register_parameter("weight", None) + + +@pytest.mark.parametrize("granularity", ["tensor", "channel"]) +def test_shared_fp8_projections_match_independent_quantization(granularity): + if not torch.cuda.is_available() or torch.cuda.get_device_capability() < (8, 9): + pytest.skip("sm89+ CUDA is required for FP8 GEMM") + from fastvideo.layers.linear import ReplicatedLinear + + torch.manual_seed(17) + layers = tuple(ReplicatedLinear(128, 256, bias=True, quant_config=FP8Config(granularity), + prefix=f"block.attn.to_{name}") for name in ("q", "k", "v")) + for layer in layers: + layer.to(device="cuda", dtype=torch.bfloat16) + layer.weight.data.normal_(std=0.1) + layer.bias.data.normal_(std=0.1) + _install_fp8_buffers(layer) + x = torch.randn(1, 272, 128, device="cuda", dtype=torch.bfloat16) + with torch.inference_mode(): + reference = [layer(x)[0] for layer in layers] + with patch.object(FP8QuantizeMethod, "quantize_input", autospec=True, + side_effect=FP8QuantizeMethod.quantize_input) as quant: + actual = _shared_input_projections(layers, x) + assert quant.call_count == 1 + for expected, output in zip(reference, actual, strict=True): + torch.testing.assert_close(output, expected, atol=0, rtol=0) + + +@pytest.mark.parametrize("kernel", ["original", "bf16", "int8"]) +@pytest.mark.parametrize("fp8", [False, True]) +@pytest.mark.parametrize("gate_active", [False, True]) +@pytest.mark.parametrize("fused_rope", [False, True]) +def test_tile_first_matches_generic_vsa_with_partial_tiles(env_overrides, distributed_setup, tmp_path, + fp8, gate_active, fused_rope, kernel): + if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported(): + pytest.skip("BF16 CUDA is required") + if fp8 and torch.cuda.get_device_capability() < (8, 9): + pytest.skip("sm89+ CUDA is required for FP8 GEMM") + from fastvideo.attention.backends.video_sparse_attn_h3 import MiniMaxH3VSAMetadataBuilder + from fastvideo.forward_context import set_forward_context + from fastvideo.models.dits.minimax_h3 import MiniMaxH3Attention + + env_overrides.enter_context(envs.FASTVIDEO_ATTENTION_BACKEND.override("VIDEO_SPARSE_ATTN_H3")) + env_overrides.enter_context(envs.override_external("FASTVIDEO_VSA_TRITON", "1")) + env_overrides.enter_context(envs.FASTVIDEO_VSA_SM100A.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_FP4.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_TILE_FIRST.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_SM89_KERNEL.override("original")) + capture = kernel == "int8" and fp8 and gate_active and fused_rope + if capture: + env_overrides.enter_context(envs.FASTVIDEO_H3_CAPTURE_QKV.override(str(tmp_path))) + torch.manual_seed(21) + attn = MiniMaxH3Attention(256, 2, 128, 1e-5, (AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3,), FP8Config("channel") if fp8 else None, + "transformer_blocks.0.attn", fuse_qknorm_rope=fused_rope) + attn.to(device="cuda", dtype=torch.bfloat16) + for parameter in attn.parameters(): + parameter.data.normal_(std=0.1) + if not gate_active: + attn.to_gate_compress.weight.data.zero_() + if fp8: + for layer in (attn.to_q, attn.to_k, attn.to_v, attn.to_out): + _install_fp8_buffers(layer) + meta = MiniMaxH3VSAMetadataBuilder().build(current_timestep=999, patch_size=(1, 1, 1), + VSA_sparsity=0.8, packed_segments=(65, 97, (4, 6, 10)), + device=torch.device("cuda"), tile_size=64) + length = meta.total_seq_length + x = torch.randn(1, length, 256, device="cuda", dtype=torch.bfloat16) + angles = torch.randn(length, 96, device="cuda") + rope = angles.cos(), angles.sin() + with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=meta): + reference = attn(x, rope, length) + attn._vsa_tile_first = True + attn.distributed_attention.attn_impl._sm89_kernel = kernel + actual = attn(x, rope, length) + # Row order can choose a different GEMM reduction; neither attention nor + # the VSA selection/padding semantics are approximated by this route. + error = (actual.float() - reference.float()).norm() / reference.float().norm() + assert error < (0.02 if fp8 else 0.005), float(error) + torch.testing.assert_close(actual, reference, rtol=0.03, atol=0.05) + if capture: + data = torch.load(tmp_path / "layer-0.pt", weights_only=True) + torch.testing.assert_close(data["vbs"], meta.variable_block_sizes.cpu(), rtol=0, atol=0) + assert data["q"].shape == (1, 2, meta.variable_block_sizes.numel() * 64, 128) + assert data["mask"].dtype == torch.bool diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py new file mode 100644 index 0000000000..165ccc41e5 --- /dev/null +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py @@ -0,0 +1,187 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json + +import torch +import torch.nn as nn +from safetensors.torch import save_file + +from fastvideo.models.vaes.minimax_h3_int8_convrot import ( + Int8ConvRotLinear, + _int8_linear_from_tensors, + dense_vae_safetensors, + overlay_minimax_h3_int8_convrot_decoder, + parse_comfy_quant_marker, + regular_hadamard, + rotate_activation, +) + + +def test_regular_hadamard_is_orthogonal() -> None: + hadamard = regular_hadamard(16, device=torch.device("cpu"), dtype=torch.float64) + identity = torch.eye(16, dtype=torch.float64) + torch.testing.assert_close(hadamard @ hadamard.T, identity, atol=1e-12, rtol=0.0) + + +def test_convrot_rotation_is_an_involution() -> None: + torch.manual_seed(0) + x = torch.randn(4, 256) + rotated = rotate_activation(x, 256) + twice = rotate_activation(rotated, 256) + torch.testing.assert_close(twice, x, atol=1e-5, rtol=1e-5) + + +def test_int8_gemm_scales_in_float32_not_fp16() -> None: + acc = torch.tensor([[100_000]], dtype=torch.int32) + x_scale = torch.tensor([[1.0 / 127.0]], dtype=torch.float16) + weight_scale = torch.tensor([[1.0 / 127.0]], dtype=torch.float16) + overflowed = acc.to(torch.float16) * x_scale * weight_scale.t() + scaled = Int8ConvRotLinear._dequant_int8_gemm(acc, x_scale, weight_scale) + assert torch.isinf(overflowed).all() + assert torch.isfinite(scaled).all() + torch.testing.assert_close(scaled, acc.float() * x_scale.float() * weight_scale.t().float()) + + +def test_int8_linear_rejects_incompatible_convrot_group() -> None: + marker = {"format": "int8_tensorwise", "convrot": True, "convrot_groupsize": 256} + weight = torch.ones(8, 8, dtype=torch.int8) + scale = torch.ones(8, 1) + try: + _int8_linear_from_tensors(weight, scale, None, marker) + except ValueError as error: + assert "group_size 256" in str(error) + assert "in_features 8" in str(error) + else: + raise AssertionError("incompatible ConvRot overlay must be rejected") + + +def test_int8_linear_matches_dequantized_matmul() -> None: + torch.manual_seed(1) + layer = Int8ConvRotLinear(256, 32, bias=True, convrot=True, group_size=256) + layer.weight.copy_(torch.randint(-8, 8, (32, 256), dtype=torch.int8)) + layer.weight_scale.copy_(torch.linspace(0.01, 0.02, 32).unsqueeze(1)) + layer.bias.data.copy_(torch.randn(32)) + x = torch.randn(3, 256) + out = layer(x) + rotated = rotate_activation(x, 256) + expected = torch.nn.functional.linear(rotated, layer.weight.float() * layer.weight_scale, layer.bias) + torch.testing.assert_close(out, expected, atol=1e-5, rtol=1e-5) + + +def test_dense_vae_safetensors_drops_convrot_overlay() -> None: + kept = dense_vae_safetensors([ + "/tmp/vae/diffusion_pytorch_model-00001-of-00003.safetensors", + "/tmp/vae/minimax_h3_video_vae_int8_convrot.safetensors", + "/tmp/vae/other.safetensors", + ]) + assert kept == [ + "/tmp/vae/diffusion_pytorch_model-00001-of-00003.safetensors", + "/tmp/vae/other.safetensors", + ] + + +def test_parse_comfy_quant_marker_reads_padded_uint8() -> None: + payload = json.dumps({"format": "int8_tensorwise", "convrot": True, "convrot_groupsize": 256}).encode() + blob = torch.zeros(72, dtype=torch.uint8) + blob[:len(payload)] = torch.tensor(list(payload), dtype=torch.uint8) + marker = parse_comfy_quant_marker(blob) + assert marker["format"] == "int8_tensorwise" + assert marker["convrot"] is True + assert marker["convrot_groupsize"] == 256 + + +class _FakeH3VAE(nn.Module): + + def __init__(self) -> None: + super().__init__() + decoder = nn.Module() + decoder.proj_in = nn.Linear(24, 8, bias=True) + decoder.register_tokens = nn.Parameter(torch.zeros(1, 4, 8), requires_grad=False) + decoder.norm_out = nn.LayerNorm(8) + decoder.proj_out = nn.Linear(8, 12, bias=True) + block = nn.Module() + block.norm1 = nn.RMSNorm(8, elementwise_affine=True) + block.norm2 = nn.RMSNorm(8, elementwise_affine=True) + block.scale1 = nn.Parameter(torch.zeros(8), requires_grad=False) + block.scale2 = nn.Parameter(torch.zeros(8), requires_grad=False) + attn = nn.Module() + attn.heads = 2 + attn.dim_head = 4 + attn.to_q = nn.Linear(8, 8, bias=True) + attn.to_k = nn.Linear(8, 8, bias=True) + attn.to_v = nn.Linear(8, 8, bias=True) + attn.to_out = nn.ModuleList([nn.Linear(8, 8, bias=True)]) + ff = nn.Module() + swiglu = nn.Module() + swiglu.proj = nn.Linear(8, 16, bias=True) + ff.net = nn.ModuleList([swiglu, nn.Dropout(0.0), nn.Linear(8, 8, bias=True)]) + block.attn = attn + block.ff = ff + decoder.transformer_blocks = nn.ModuleList([block]) + self.decoder = decoder + + +def _marker_tensor() -> torch.Tensor: + payload = json.dumps({"format": "int8_tensorwise", "convrot": False}).encode() + blob = torch.zeros(64, dtype=torch.uint8) + blob[:len(payload)] = torch.tensor(list(payload), dtype=torch.uint8) + return blob + + +def test_overlay_splits_fused_qkv_and_ffn(tmp_path) -> None: + tensors = { + "decoder.x_embedder.weight": torch.ones(8, 24), + "decoder.x_embedder.bias": torch.zeros(8), + "decoder.register_tokens": torch.ones(1, 4, 8), + "decoder.norm_out.weight": torch.ones(8), + "decoder.norm_out.bias": torch.zeros(8), + "decoder.proj_out.weight": torch.ones(12, 8), + "decoder.proj_out.bias": torch.zeros(12), + "decoder.transformer_blocks.0.attn.to_qkv.weight": torch.arange(24 * 8, dtype=torch.int8).reshape(24, 8), + "decoder.transformer_blocks.0.attn.to_qkv.weight_scale": torch.ones(24, 1), + "decoder.transformer_blocks.0.attn.to_qkv.bias": torch.zeros(24), + "decoder.transformer_blocks.0.attn.to_qkv.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.attn.to_out.weight": torch.ones(8, 8, dtype=torch.int8), + "decoder.transformer_blocks.0.attn.to_out.weight_scale": torch.ones(8, 1), + "decoder.transformer_blocks.0.attn.to_out.bias": torch.zeros(8), + "decoder.transformer_blocks.0.attn.to_out.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.ff.w1.weight": torch.cat( + [torch.ones(8, 8, dtype=torch.int8), torch.full((8, 8), 2, dtype=torch.int8)], dim=0), + "decoder.transformer_blocks.0.ff.w1.weight_scale": torch.cat( + [torch.ones(8, 1), torch.full((8, 1), 0.5)], dim=0), + "decoder.transformer_blocks.0.ff.w1.bias": torch.cat([torch.ones(8), torch.full((8, ), 3.0)], dim=0), + "decoder.transformer_blocks.0.ff.w1.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.ff.w2.weight": torch.ones(8, 8, dtype=torch.int8), + "decoder.transformer_blocks.0.ff.w2.weight_scale": torch.ones(8, 1), + "decoder.transformer_blocks.0.ff.w2.bias": torch.zeros(8), + "decoder.transformer_blocks.0.ff.w2.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.norm1.weight": torch.ones(8), + "decoder.transformer_blocks.0.norm2.weight": torch.ones(8), + "decoder.transformer_blocks.0.scale1": torch.ones(8), + "decoder.transformer_blocks.0.scale2": torch.ones(8), + } + path = tmp_path / "minimax_h3_video_vae_int8_convrot.safetensors" + save_file(tensors, path) + vae = _FakeH3VAE() + installed = overlay_minimax_h3_int8_convrot_decoder(vae, path) + assert installed == 6 + q = vae.decoder.transformer_blocks[0].attn.to_q + k = vae.decoder.transformer_blocks[0].attn.to_k + v = vae.decoder.transformer_blocks[0].attn.to_v + fused = tensors["decoder.transformer_blocks.0.attn.to_qkv.weight"].view(2, 3, 4, 8) + assert isinstance(q, Int8ConvRotLinear) + assert q.weight.shape == (8, 8) + assert torch.equal(q.weight, fused[:, 0].reshape(8, 8)) + assert torch.equal(k.weight, fused[:, 1].reshape(8, 8)) + assert torch.equal(v.weight, fused[:, 2].reshape(8, 8)) + proj = vae.decoder.transformer_blocks[0].ff.net[0].proj + assert isinstance(proj, Int8ConvRotLinear) + assert isinstance(vae.decoder.transformer_blocks[0].ff.net[2], Int8ConvRotLinear) + assert torch.equal(proj.weight[:8], torch.full((8, 8), 2, dtype=torch.int8)) + assert torch.equal(proj.weight[8:], torch.ones(8, 8, dtype=torch.int8)) + assert torch.equal(proj.weight_scale[:8], torch.full((8, 1), 0.5)) + assert torch.equal(proj.bias[:8], torch.full((8, ), 3.0)) + assert not hasattr(vae.decoder, "mask_token") + assert torch.equal(vae.decoder.proj_in.weight, torch.ones(8, 24)) diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py new file mode 100644 index 0000000000..87caab3ee3 --- /dev/null +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py @@ -0,0 +1,95 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exact parity of shared VAE input preparation and strided INT8 weight GEMMs.""" +from unittest.mock import patch + +import pytest +import torch + +import fastvideo.envs as envs + +from fastvideo.models.vaes.minimax_h3_int8_convrot import Int8ConvRotLinear, shared_int8_projections + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for INT8 GEMM") +@pytest.mark.parametrize("rows", [3, 17, 129]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("convrot", [False, True]) +def test_shared_int8_and_transpose_views_are_exact(rows, dtype, convrot, env_overrides): + env_overrides.enter_context(envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.override(False)) + torch.manual_seed(73) + layers = tuple(Int8ConvRotLinear(256, out, bias=index != 1, convrot=convrot, group_size=256) + .to("cuda") for index, out in enumerate([128, 256, 64])) + for layer in layers: + layer.weight.random_(-127, 128) + layer.weight_scale.uniform_(0.0001, 0.03) + if layer.bias is not None: + layer.bias.data.normal_() + x = torch.randn(1, rows, 256, device="cuda", dtype=dtype) + x[0, 0].zero_() # clamp/padding semantics must also survive sharing + with torch.inference_mode(): + expected = tuple(layer(x) for layer in layers) + with patch.object(Int8ConvRotLinear, "quantize_input", autospec=True, + side_effect=Int8ConvRotLinear.quantize_input) as quant: + shared = shared_int8_projections(layers, x) + assert quant.call_count == 1 + for layer in layers: + layer._transpose_view = True + layer._fused_dequant = True + views = shared_int8_projections(layers, x) + for ref, actual, view in zip(expected, shared, views, strict=True): + assert torch.isfinite(ref).all() + torch.testing.assert_close(actual, ref, rtol=0, atol=0) + torch.testing.assert_close(view, ref, rtol=0, atol=0) + + +def test_shared_int8_keeps_cpu_fallback_exact(): + layers = tuple(Int8ConvRotLinear(16, 8, bias=False, convrot=False, group_size=16) for _ in range(3)) + for layer in layers: + layer.weight.fill_(1) + layer.weight_scale.fill_(0.01) + x = torch.ones(3, 16) + for actual, ref in zip(shared_int8_projections(layers, x), (layer(x) for layer in layers), strict=True): + torch.testing.assert_close(actual, ref, rtol=0, atol=0) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for VAE attention parity") +def test_vae_attention_shares_quantized_projections_exactly(distributed_setup, env_overrides): + from fastvideo.models.vaes.minimax_h3_video import MiniMaxH3VideoAttention + + env_overrides.enter_context(envs.FASTVIDEO_H3_VAE_INT8_SHARED_QKV.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.override(False)) + torch.manual_seed(49) + attention = MiniMaxH3VideoAttention(256, 2, 128).to("cuda").eval() + for name in ("to_q", "to_k", "to_v"): + layer = Int8ConvRotLinear(256, 256, bias=True, convrot=True, group_size=256).to("cuda") + layer.weight.random_(-8, 9) + layer.weight_scale.fill_(0.01) + layer.bias.data.normal_(std=0.1) + setattr(attention, name, layer) + x = torch.randn(2, 33, 256, device="cuda") + with torch.inference_mode(): + expected = attention(x) + attention._share_int8_qkv = True + for layer in (attention.to_q, attention.to_k, attention.to_v): + layer._transpose_view = True + layer._fused_dequant = True + actual = attention(x) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for fused INT8 epilogue") +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("has_bias", [False, True]) +def test_fused_int8_epilogue_large_accumulators_and_small_scales(dtype, has_bias): + from fastvideo.models.vaes.minimax_h3_int8_kernels import fused_int8_dequant_bias + + torch.manual_seed(19) + acc = torch.randint(-400_000_000, 400_000_000, (129, 264), device="cuda", dtype=torch.int32) + x_scale = torch.logspace(-30, -3, 129, device="cuda").view(-1, 1) + w_scale = torch.logspace(-6, -2, 264, device="cuda").view(-1, 1) + bias = torch.randn(264, device="cuda") if has_bias else None + expected = acc.float() * x_scale.float() * w_scale.t().float() + if bias is not None: + expected = expected + bias.float() + actual = fused_int8_dequant_bias(acc, x_scale, w_scale, bias, dtype) + torch.testing.assert_close(actual, expected.to(dtype), rtol=0, atol=0) diff --git a/fastvideo/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index d0712091c6..f4de6a30f0 100644 --- a/fastvideo/worker/gpu_worker.py +++ b/fastvideo/worker/gpu_worker.py @@ -21,6 +21,38 @@ def _log_cuda_device_uuid(rank: int, device: torch.device) -> None: logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False) +def _log_pipeline_memory(pipeline) -> None: + """Debug (FASTVIDEO_MEMORY_REPORT=1): bytes held per pipeline component, by device and dtype, plus the + largest tensors, so the resident footprint can be attributed before choosing offload placements.""" + gib = 1024**3 + for name, module in getattr(pipeline, "modules", {}).items(): + if not isinstance(module, torch.nn.Module): + continue + by_kind: dict[str, int] = {} + largest: list[tuple[int, str, str]] = [] + seen: set[int] = set() + for tname, t in list(module.named_parameters()) + list(module.named_buffers()): + if t is None or id(t) in seen: + continue + seen.add(id(t)) + nbytes = t.numel() * t.element_size() + key = f"{t.device.type}/{str(t.dtype).replace('torch.', '')}" + by_kind[key] = by_kind.get(key, 0) + nbytes + largest.append((nbytes, tname, key)) + largest.sort(reverse=True) + total = sum(by_kind.values()) + logger.info("MEMREPORT %s total=%.2f GiB %s", name, total / gib, { + k: round(v / gib, 2) + for k, v in sorted(by_kind.items(), key=lambda kv: -kv[1]) + }) + for nbytes, tname, key in largest[:8]: + logger.info("MEMREPORT %s %.3f GiB %s %s", name, nbytes / gib, key, tname) + if torch.cuda.is_available(): + logger.info("MEMREPORT cuda allocated=%.2f GiB reserved=%.2f GiB", + torch.cuda.memory_allocated() / gib, + torch.cuda.memory_reserved() / gib) + + class Worker: def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str): @@ -71,6 +103,12 @@ def init_device(self) -> None: # Set the CUDA device BEFORE any CUDA calls if current_platform.is_cuda_alike(): torch.cuda.set_device(self.device) + # Debug: FASTVIDEO_CUDA_MEMORY_CAP_GIB emulates a smaller card by capping this process's allocator. + cap_gib = envs.FASTVIDEO_CUDA_MEMORY_CAP_GIB.get() + if cap_gib > 0: + total = torch.cuda.get_device_properties(self.device).total_memory + torch.cuda.set_per_process_memory_fraction(min(1.0, cap_gib * 1024**3 / total), self.device) + logger.info("Capped CUDA allocator at %s GiB of %.1f GiB", cap_gib, total / 1024**3) 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) @@ -90,6 +128,8 @@ def init_device(self) -> None: self.distributed_init_method) self.pipeline = build_pipeline(self.fastvideo_args) + if envs.FASTVIDEO_MEMORY_REPORT.get() and self.rank == 0: + _log_pipeline_memory(self.pipeline) def execute_forward(self, forward_batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args) diff --git a/mkdocs.yml b/mkdocs.yml index 81c0840163..f3e3e45f30 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -183,6 +183,7 @@ nav: - Configuration: inference/configuration.md - Offloading: inference/offloading.md - Optimizations: inference/optimizations.md + - FastH3 on RTX PRO 6000: inference/fasth3_rtx_pro_6000.md - ComfyUI: inference/comfyui.md - Support Matrix: inference/support_matrix.md - CLI: inference/cli.md diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md new file mode 100644 index 0000000000..4b7132f028 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -0,0 +1,458 @@ +# FastH3 on a single RTX 4090 + +Use the private `FastVideo/FastH3-Pruned-8Step-FP8-ckpt300` checkpoint. +Keep its `fastvideo_inference.json`: nine sigma-grid points produce eight +DMD forwards. Preserve VSA sparsity 0.8 and tile size 64. + +## Setup and validation + +The consumer launch uses seed 1234 and 124 frames for both resolutions. +Pass `--seed 1234 --frames 124 --timed 4` with the two launch benchmark +prompts to collect one warmup and two timed samples per prompt. + +On hosts with limited RAM, `--pageable-host` sets the existing engine +`offload.pin_cpu_memory` option to false. Layerwise DiT and encoder offload +then retain existing CPU storage instead of copying it into pinned arenas. +Checkpoint-backed pages can be reclaimed by the OS, at the cost of slower +host-to-device transfers. This also allows a cached-component recipe to be +tested before resorting to `--lazy`, which reloads components per request. +Validate actual host and GPU peaks for each recipe; disabling pinning does +not make ordinary heap allocations reclaimable checkpoint pages. + +The memory sampler supports cgroup v2 (`memory.current`, `anon`) and v1 +(`memory.usage_in_bytes`, anonymous RSS). Both report pod-wide usage including +file cache separately from anonymous memory. If host counters cannot be read, +the host peak is unknown and NVML GPU capacity sampling continues. +Failed generate calls retain sampled memory and an error in `failed_runs`, +separately from completed timing samples. + +Full-rank V2 can replace its large AdaLN projection weights with exact outputs +for the checkpoint's fixed T2AV DMD ladder: + +```bash +python -P scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py \ + /workspace/vol/v2-fp8 /workspace/v2-adaln.pt --source-commit +export FASTVIDEO_H3_ADALN_TABLE=/workspace/v2-adaln.pt +``` + +The helper checks every block/rung against the original modulation module, +writes the exact inputs alongside the table, and records hashes and checkpoint +provenance in a JSON sidecar. Validate the real pipeline timestep embeddings +against the saved inputs before using a new table in headline measurements. +Use the table with its matching checkpoint and T2AV contract; unsupported +timestep keys raise rather than silently changing the schedule. + +The October 3, 2026 pod has one RTX 4090 (24,564 MiB), driver 580.126.20, +a 99,999,997,952-byte host cgroup limit, and 150 GB disk. Its runtime is +PyTorch 2.12.0+cu126, CUDA toolkit 12.6, and FlashInfer 0.7.1rc2. + +Exact-size host arenas replace the pinned allocator for layerwise blocks +and H3 module swaps. Each arena packs typed views at 256-byte offsets into +dedicated CUDA-registered pages. Live views retain their registration owner. +Mutation and hook detachment unregister old arenas. Registration failures +fall back to the ordinary pinned allocator with a warning. + +Validation on this pod: all 12 tests passed with the command below. +The tests cover repeated offloaded forwards, BF16, mixed-dtype exact copies, +owner lifetime, registration fallback, large-buffer mutation, detachment +after prefetch, and persistent H3 swaps with changing buffers. + +```bash +source /workspace/env.sh +source /workspace/venv/bin/activate +cd /workspace/fastvideo +python -P -m pytest fastvideo/tests/hooks/test_pinned_memory.py \ + fastvideo/tests/hooks/test_layerwise_offload.py -q +``` + +The supplied `handoff_4090/kernel_microbench/pinned_memory.py` measured +5.06 GiB extra cgroup usage for 2.87 GiB through `pin_memory()`, versus +2.87 GiB using direct host registration. Pinned H2D measured 25.9 GB/s; +pageable H2D measured 10.2 GB/s. These are microbenchmarks, not clip timings. + +The supplied `handoff_4090/kernel_microbench/t_fp8.py` measured: + +| K → N, M = 38,976 | Fused quantization | FP8 GEMM + scale epilogue | +| --- | --- | --- | +| 5376 → 5376 | 0.68 ms | 11.27 ms | +| 5376 → 28672 | 0.68 ms | 40.39 ms | +| 14336 → 5376 | 2.10 ms | 27.96 ms | + +Commands on the pod, run before clip benchmarks: + +```bash +cd /workspace +python -P kernel_microbench/pinned_memory.py +python -P kernel_microbench/t_fp8.py +``` + +## End-to-end baseline + +Run one warmup and at least two timed requests. The benchmark saves clips, +the exact Python command, runtime environment, sampling geometry, config, +source commit, each wall time, and the median to `results.json`. Stage logs +include GPU allocation peaks and conditioning, denoise, and decode timings. +Per-run host samples report both total cgroup memory and anonymous memory; +total usage includes checkpoint file cache. Both are sampled every 100 ms +and include all processes in the pod's cgroup. + +```bash +cd /workspace +export FASTVIDEO_SOURCE_COMMIT=ce877b200 +export FASTVIDEO_H3_PARK_MODULES=vae,audio_vae +export MAX_JOBS=4 +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + baseline-480p /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --lazy --no-vae-compile \ + --height 480 --width 832 --frames 243 --timed 2 +``` + +For the full-resolution baseline, change the name to `baseline-768p`, +height to 768, and width to 1344; retain 243 frames. Layerwise offload, +lazy component loading, and eager VAE decode are the starting configuration. +Do not compare these numbers with a different frame count or decoder. + +Completed baseline at source commit `ce877b200`, checkpoint revision +`f2ef54f9ff2091762ab8689b6514dcab5bc1d383`: + +| Configuration | Median e2e | Denoise stage | Video decode stage | Peak GPU allocated | Peak host anon | Peak total cgroup | +| --- | --- | --- | --- | --- | --- | --- | +| FP8, layerwise, lazy, eager H3 VAE, 832×480, 243 frames | 163.34 s | 91.72 s | 34.54 s | 17.47 GiB | 28.91 GiB | 76.95 GiB | + +Two timed requests took 163.67 s and 163.00 s after one warmup. +Stage times are medians and include deferred loading. Memory columns are +maxima across the timed requests. Total cgroup usage includes file cache. +The GPU peak occurs during conditioning. The saved clip contains 243 frames +at 24 fps (10.125 seconds) and an AAC audio track. A contact-sheet inspection +confirms a coherent pottery scene; speech and same-seed reference parity +still need review before claiming quality equivalence. + +Summarize a completed run while excluding warmup: + +```bash +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/summarize.py \ + /workspace/outputs/baseline-480p/results.json /workspace/baseline-480p.log +``` + +After a baseline works, measure FFN chunk sizes 16,384 and 8,192, then +increase resident DiT blocks within the measured GPU budget. Kernel or +decoder changes also need same-seed visual and auditory comparison. + +## Tile-first attention and full-resolution profiling + +Commit `c5d9f8132` shares compatible FP8 Q/K/V activation quantization and +releases dead block activations before residual modulation. The opt-in +`FASTVIDEO_H3_VSA_TILE_FIRST=1` scatters the attention input before projection, +then uses the existing BF16 VSA kernel. It retains tile-64 selection, partial +key validity and the learned compression branch. It supports eager, +single-rank inference; grad, compile, and multi-rank requests use the generic +path. Ten CUDA tests passed on the 4090, including mixed partial tiles, +active/zero gates, fused/unfused RoPE and FP8/nonquantized projections. + +The original 1344×768 baseline failed with a GPU OOM in post-attention +modulation before the FFN. The tile-first/FFN-16384 profiling run subsequently completed all three +768p clips without OOM: diagnostic median 349.76 s, denoise stage 244.07 s, +video decode stage 65.94 s, peak GPU allocated 17.55 GiB. A run without +profiling is still required for a release speed claim. +The profiling run used the following command. For SSH on macOS, +`-o UseKeychain=yes` retrieves the stored passphrase when the agent has no +loaded identities. Profiling/capture timings are +for diagnosis and must not be used as the final speed claim. + +```bash +FASTVIDEO_SOURCE_COMMIT=c5d9f8132 \ +FASTVIDEO_H3_PARK_MODULES=vae,audio_vae \ +FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 \ +FASTVIDEO_H3_VSA_TILE_FIRST=1 \ +FASTVIDEO_H3_CAPTURE_QKV=/workspace/qkv-768p MAX_JOBS=4 \ +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + tile-first-768p-profile /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --lazy --no-vae-compile \ + --height 768 --width 1344 --frames 243 --timed 2 --profile +``` + +`FASTVIDEO_H3_CAPTURE_QKV` saves the first inputs from layers 0, 20 and 41, +two heads each, with full real sequences, masks, valid tile sizes and packed +row indices. Disable both capture and profiling for final clip timings. + +The separate `minimax_h3_sparse_int8.py` prototype uses INT8 QK and FP8 PV, +FP32 accumulators and the original 64-token mask. It has no automatic pipeline +route. Initial real-QKV tests found approximately 1.9× fine-kernel speedup +but 4.3–13.4% relative L2 error and a strict partial-tile elementwise test +failure. Do not select it for shipping. The microbenchmark now includes +BF16 QK/PV ablations to isolate that error. Offline compilation with Triton 3.8.0 for sm89 passed all four entry points +and emitted native INT8 and FP8 MMA instructions (20,480 bytes of shared +memory for the attention kernel). This is compilation evidence only. It must +pass CUDA tests and real-QKV/clip checks before integration. +Run its microbenchmark on an idle GPU: + +```bash +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py \ + /workspace/qkv-768p --output /workspace/sparse-qkv-results.json +``` + +SpargeAttn at `ae5b629ebb41e41f86b3ea2ab5a3283f13ac151a` built on the pod +with CUDA 12.8, `TORCH_CUDA_ARCH_LIST=8.9`, and `MAX_JOBS=4`. The upstream +`-Xcompiler -include,cassert` workaround was removed from `setup.py` to +avoid GCC 13 duplicate standard-library definitions. It is not selected by +the pipeline: its public 128-query/64-key adapter also needs correct masking +of partial H3 tiles before a meaningful parity comparison. + +## Cached-component 480p result + +At `4d9846573`, retain components between requests (omit `--lazy`), and set +`FASTVIDEO_H3_VSA_TILE_FIRST=1` and `FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384`. +Keep all other baseline settings, including the eager light VAE and original +BF16 attention kernel. After one warmup the two timed requests took 115.87 s +and 115.03 s, median **115.45 s** (29.3% less wall time than the lazy baseline). +Conditioning/denoise/video-decode stage medians were 10.85/72.62/24.98 s. +Peak GPU allocation was 19.02 GiB, host anon 43.96 GiB, total cgroup 92.72 GiB +including file cache. This recipe requires more host RAM than the 32 GB target; +its minimum RAM has not been tested under a smaller host limit. + +Decoded raw video and PCM audio SHA256 hashes match the baseline exactly for +both ceramics and harbor at seed 20260929. This establishes output identity +for these two prompts; it does not establish the checkpoint's BF16-reference +quality on other prompts. Raw results, clips and hash evidence are saved in +`output/fasth3-4090-20261003/` beside the workspace. + +## sm89 kernel precision choices + +`FASTVIDEO_H3_VSA_SM89_KERNEL=bf16` opts into the new entirely BF16 tile-64 +fine kernel. `int8` uses per-token INT8 QK with BF16 PV. `original` is the +unchanged default. Resolution happens when the backend is constructed; +unsupported devices, grad and compile requests retain the original route. +Both preserve the original tile selection, partial key masks and gated +compression. The rejected FP8-PV experiment is only exposed in the diagnostic +microbenchmark, never the pipeline route. + +Two-head real-QKV captures at 1344×768, layers 0/20/41, measured: + +| QK / PV | Fine-kernel speedup including input quantization | Relative L2 vs original BF16 | +| --- | --- | --- | +| BF16 / BF16 | 1.23× | 0.005–0.008% | +| INT8 / BF16 | 1.59× | 0.58–0.62% | +| INT8 / FP8 (rejected) | 1.91× | 4.3–13.4% | + +These are fine-kernel microbenchmarks, not end-to-end clip speedups. Same-seed +clip checks are required for the INT8 route. All 44 targeted CUDA/CPU checks +passed for native/tile-first routing, partial tiles, learned compression, +shared FP8 projections and sequential component restoration. + +At `9a8465ac4`, CPU-offloaded VAEs also remain on the host during denoising; +the encode/decode stages move them on demand. This frees room for resident +DiT blocks without changing any model arithmetic. + +## Six resident blocks and encoder priorities + +At `6d3c4cda5`, the opt-in BF16 fine kernel with six resident DiT blocks, +cached components and the settings below measured **111.27 s** median for +832×480, 243 frames. Timed requests were 111.61/110.92 s after one warmup. +Conditioning, denoise and video decode medians were 11.20/67.56/25.49 s. +Peak GPU allocation was 21.63 GiB, host anonymous memory 42.10 GiB and +total cgroup usage 90.94 GiB, including file cache. + +```bash +FASTVIDEO_SOURCE_COMMIT=6d3c4cda5 \ +FASTVIDEO_H3_PARK_MODULES=vae,audio_vae \ +FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 \ +FASTVIDEO_H3_VSA_TILE_FIRST=1 FASTVIDEO_H3_VSA_SM89_KERNEL=bf16 \ +FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=6 MAX_JOBS=4 \ +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + sm89-bf16-480p-resident6 /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --no-vae-compile \ + --height 480 --width 832 --frames 243 --timed 2 +``` + +The BF16 kernel changes floating-point reductions: the ceramics clip is +visually coherent in the sampled contact sheet but is not identical to the +original-kernel clip (decoded-video SSIM 0.597653). This is a speed candidate, +not proof of quality equivalence. The original kernel remains the default. + +The same cached/six-resident-block BF16 recipe at `26390848b`, with name +`sm89-bf16-480p-5s-resident6` and `--frames 124`, measured **65.27 s** median. +The legal frame count represents 5.167 s at 24 fps. Timed requests were +65.08/65.46 s after a 110.34 s warmup. Conditioning/denoise/video-decode +medians were 11.74/34.22/13.44 s. Peak GPU allocation was 21.62 GiB, +host anon 41.14 GiB and total cgroup 88.45 GiB. These generation wall times +include decode/export and exclude initial generator construction. No +profiling or QKV capture was enabled. + +At `26390848b`, per-key-tile V scaling reduced experimental INT8-QK/FP8-PV +real-tensor error to 0.84–1.54%, with 1.77× fine-kernel speedup. Dynamic +per-query, per-key-block P scaling also preserves contributions that would +underflow with the fixed P scale; it measured 0.80–1.52% error and 1.71× +speedup. All 30 focused CUDA kernel/routing tests passed. These FP8-PV routes +remain microbenchmark-only and need clip validation. + +The current text encoder is the trimmed 50-layer Qwen3-VL with serialized +NVFP4 weights, dequantized to BF16 per linear on sm89. The existing serialized +blockwise FP8 encoder requires sm100+ and FlashInfer's Blackwell GEMM; it +cannot run on the 4090 as written. An Ada FP8 implementation would also need +encoder streaming because its weights are larger. Reading the current +checkpoint tensor shapes gives 15.33 GiB total encoder weights, including +11.35 GiB packed values and 1.42 GiB block scales. Replacing those packed +values with FP8 while retaining the other tensors projects about 25.3 GiB +before activations (the FP8 block-scale overhead is small). This is a storage +estimate, not a measured FP8 encoder. First try fused NVFP4 +dequantization and avoid per-linear GPU scalar synchronization; then compare +a native sm89 FP8 encoder at equal prompts. The later streamed/fused conditioning stage is about 0.55 s, so a new +encoder export must be measured against that implementation. + +Remaining speed experiments include VAE compilation and tile-batch tuning, +INT8 decoder epilogue fusion, direct strided fine-attention reads, and fused +norm/activation quantization. See the completed smaller-VRAM results below. +The cached recipe's host peak does not establish a 32 GB system-RAM minimum. + + +## Streamed encoder and smaller VRAM caps + +At `753e560f6`, `FASTVIDEO_H3_ENCODER_LAYERWISE=1` streams the language +layers separately from DiT residency, retaining token embeddings and unused +vision modules on the CPU. This route currently supports text-only T2VA; +visual references fail explicitly. The pipeline preserves the streamed +placement. Twenty-eight offload/encoder/stage tests passed, including exact +repeated BF16 and NVFP4 parity. + +At `78540b635`, `FASTVIDEO_H3_ENCODER_FUSED_DEQUANT=1` expands packed NVFP4 +weights with one Triton pass. Fifteen strict tests passed across FP32, FP16, +BF16, swizzled scales and repeated encoder forwards. On three large stress +matrices the expansion was 25.3–25.9× faster and used 8× less temporary GPU +memory than Torch expansion. This is a dequantization microbenchmark; the +whole conditioning stage measured 0.55–0.62 seconds in the clip runs below. +The current encoder remains NVFP4 storage with BF16 GEMMs on Ada. + +All rows use 832×480, 243 frames, eight DMD forwards, sparsity 0.8, tile 64, +cached components, INT8 QK/BF16 PV and the eager light H3 VAE. Each median +has one warmup and two timed requests. Source is `78540b635` except the +16 GiB row (`753e560f6`, before fused dequantization). + +| 4090 configuration | Median e2e | Timed requests | Denoise | Video decode | Peak GPU allocated | Peak host anon | +| --- | --- | --- | --- | --- | --- | --- | +| 16 GiB cap, 6 resident | 104.03 s | 100.10 / 107.96 s | 72.42 s | 25.54 s | 11.28 GiB | 39.86 GiB | +| 12 GiB cap, 0 resident | 107.27 s | 111.02 / 103.52 s | 77.38 s | 25.48 s | 8.68 GiB | 42.35 GiB | +| Uncapped, 30 resident | 98.97 s | 98.88 / 99.05 s | 69.12 s | 25.47 s | 21.69 GiB | 29.44 GiB | + +Set `FASTVIDEO_CUDA_MEMORY_CAP_GIB=12` for the 12 GiB recipe; unset it for +the full card. Set resident blocks to the table value. Both fused rows use: + +```bash +FASTVIDEO_SOURCE_COMMIT=78540b635 \ +FASTVIDEO_H3_PARK_MODULES=vae,audio_vae \ +FASTVIDEO_H3_ENCODER_LAYERWISE=1 FASTVIDEO_H3_ENCODER_FUSED_DEQUANT=1 \ +FASTVIDEO_H3_VSA_TILE_FIRST=1 FASTVIDEO_H3_VSA_SM89_KERNEL=int8 \ +FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=30 FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 \ +FASTVIDEO_H3_VAE_TILE_BATCH=28 MAX_JOBS=4 \ +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + sm89-int8-480p-resident30-fused /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --no-vae-compile --height 480 --width 832 --frames 243 --timed 2 +``` + +The 12 GiB and 30-resident ceramics clips have identical decoded-video and +PCM audio hashes to the six-resident INT8 clip. These placement and dequant +changes preserve the candidate's output; quality equivalence of INT8 +attention to the original BF16 checkpoint still requires motion/speech review. +Allocator caps emulate available VRAM on a 4090, not another card's speed. +Host peaks include all pod processes; actual 32 GB host-limit support has +not been established. The 30-resident warmup reached 30.26 GiB anonymous +memory and the timed runs reached 77.50 GiB total cgroup usage including cache. + +The 8 GiB cap at `fcdba37fc` completed its warmup but OOMed on the timed +harbor prompt in fine attention. Do not report it as supported. A 34-resident +experiment completed denoising but OOMed during VAE INT8 epilogue allocation. +Both failures motivate subsequent memory work rather than speed claims. + + +## Consumer kernel memory work on the release core + +The consumer branch is rebased onto release core `a97d23f09` (fork PR #45). +Historical measured commits remain reachable through tag +`h3-consumer-fp8-measured-20261003`; rebase changes their branch commit IDs. + +`7a0d7d33b` lets the INT8-QK/BF16-PV fine kernel read BSHD-backed views +without retaining three full BF16 layout copies. Q/K quantization writes +contiguous INT8 arrays and the output stays BHSD; the K-mean reduction keeps +the reference's contiguous reduction order. Four CUDA tests require exact +output equality for partial tiles, empty selections, multiple batches and +partner padding, together with a lower peak allocation. + +`3c0668f6c` adds the opt-in eager +`FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT=1`. One Triton pass applies the INT32 +GEMM's row scale, output-channel scale and optional bias, then casts the +result. FP32 operations retain separate rounding steps (FP fusion disabled). +Strict tests cover zero rows, small input batches, FP32/FP16/BF16, bias, +large INT32 accumulators and tiny scales. Compiled and grad paths retain +the reference implementation. Combine it with shared QKV and transpose +views using `FASTVIDEO_H3_VAE_INT8_SHARED_QKV=1` and +`FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW=1`. + +After the core rebase, 120 focused tests passed (one unrelated GPU cudagraph +check excluded) and pre-commit passed. `87b22a5d8` adapts capture and tests +to packed-segment metadata while retaining the core's calibrated NVFP4 +activation-scale guard. `fb92af176` adds optional NVML total-device-memory +samples every 100 ms, alongside host samples. Summaries distinguish sampled +total GPU usage from PyTorch's allocated peak. Allocator caps omit driver +and external CUDA memory, so the 8 GiB total-budget experiment uses a +7.25 GiB allocator cap and must also satisfy the observed NVML budget. +That trial completed its warmup and first timed request, but sampled total +GPU usage reached about 8.28 GiB. It therefore does **not** meet a strict +8 GiB device target. A tighter allocator cap still needs validation. +These are 4090 simulations; real lower-VRAM and 30-series performance still +needs those devices. + + +## Warmed release-core clip results + +At `fb92af176`, after one warmup and two timed requests: + +| Clip | Median e2e | Timed requests | Conditioning | Denoise | Video decode | Peak allocated | Sampled total GPU | Peak host anon | +| --- | --- | --- | --- | --- | --- | --- | --- | --- | +| 832×480, 124 frames / 5.167 s | **41.75 s** | 42.16 / 41.33 s | 0.55 s | 32.69 s | 6.63 s | 22.01 GiB | 23.983 GiB | 24.78 GiB | +| 832×480, 243 frames / 10.125 s | **79.67 s** | 79.64 / 79.70 s | 0.55 s | 63.51 s | 13.23 s | 20.29 GiB | 23.985 GiB | 27.02 GiB | + +Warmups took 84.40 / 125.16 s. These wall times include audio, frame export +and MP4 saving, and exclude initial generator construction. The 5 s recipe +keeps 34 DiT blocks resident; the 10 s recipe keeps 30. Both use the same +checkpoint revision and eight DMD forwards as the earlier rows. No profiling, +QKV capture, alternate decoder or sparsity increase is enabled. + +```bash +source /workspace/env.sh +source /workspace/venv/bin/activate +cd /workspace +export PYTHONPATH="/workspace/fastvideo-core:${PYTHONPATH:-}" +export FASTVIDEO_SOURCE_COMMIT=fb92af176 +export FASTVIDEO_H3_PARK_MODULES=vae,audio_vae +export FASTVIDEO_H3_ENCODER_LAYERWISE=1 FASTVIDEO_H3_ENCODER_FUSED_DEQUANT=1 +export FASTVIDEO_H3_VSA_TILE_FIRST=1 FASTVIDEO_H3_VSA_SM89_KERNEL=int8 +export FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=34 FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 +export FASTVIDEO_H3_VAE_TILE_BATCH=28 +export FASTVIDEO_H3_VAE_INT8_SHARED_QKV=1 FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW=1 +export FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT=1 MAX_JOBS=4 +unset FASTVIDEO_CUDA_MEMORY_CAP_GIB +python -P /workspace/fastvideo-core/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + sm89-int8-fast2-480p-5s /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --no-vae-compile --height 480 --width 832 --frames 124 --timed 2 +``` + +For 10 s, set resident blocks to 30, name to `sm89-int8-fast2-480p-10s`, +and frames to 243. Both decoded video and PCM audio hash-identically to the +previous INT8 candidate for ceramics and harbor at 10 s. This validates these +memory/decode changes on those prompts, while the INT8 attention candidate +still differs from the original BF16 attention clips and needs full quality +review. The sampled 5 s contact sheet is coherent. Raw clips, hashes and +results live in `output/fasth3-4090-20261003/` beside the worktree. + +The earlier unprofiled 1344×768, 243-frame run at `fcdba37fc` completed at +279.94 s median (289.06 / 270.82 s, 306.80 s warmup), with 12 resident blocks +and shared VAE QKV/transpose views, before direct-layout attention and fused +VAE epilogues. Its stage medians were 0.55 s conditioning, 227.21 s denoise, +44.82 s video decode and 0.48 s audio. Peak allocation was 21.19 GiB and host +anonymous memory 39.81 GiB. The updated 768p run was queued after the +7.25 GiB allocator-cap trial. SSH became unreachable before the final +results could be collected; updated 768p speed and strict 8 GiB support +remain unverified. + +Track B is staged in [draft PR #46](https://github.com/aryan5v/FastVideo/pull/46), +stacked on the shared release core in #45. Historical timing sources are +preserved by the `h3-consumer-fp8-measured-20261003` tag. diff --git a/scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py b/scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py new file mode 100644 index 0000000000..798793e678 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py @@ -0,0 +1,59 @@ +"""Compare serialized NVFP4 weight expansion on an idle consumer GPU.""" +import argparse +import json +import os +import pathlib +import time + +import torch + +from fastvideo.layers.quantization.nvfp4_dequant import dequantize_nvfp4_cuda +from fastvideo.models.encoders.minimax_h3_checkpoint_nvfp4 import dequantize_serialized_nvfp4 + + +def measure(fn, packed, scales, global_scale): + for _ in range(3): + fn(packed, scales, global_scale) + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + start, stop = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + wall = time.perf_counter() + start.record() + for _ in range(10): + fn(packed, scales, global_scale) + stop.record() + stop.synchronize() + return {"gpu_ms": start.elapsed_time(stop) / 10, + "wall_ms": (time.perf_counter() - wall) * 100, + "peak_extra_gib": (torch.cuda.max_memory_allocated() - baseline) / 2**30} + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--output", type=pathlib.Path, required=True) + args = ap.parse_args() + rows = [] + global_scale = float(torch.tensor(317.224, dtype=torch.float32)) + torch.manual_seed(63) + for n, k in [(8192, 8192), (25600, 8192), (8192, 25600)]: + packed = torch.randint(0, 256, (n, k // 2), device="cuda", dtype=torch.uint8) + scales = torch.randint(0, 127, (n, k // 16), device="cuda", dtype=torch.uint8) + reference = dequantize_serialized_nvfp4(packed, scales, global_scale) + fused = dequantize_nvfp4_cuda(packed, scales, global_scale) + torch.testing.assert_close(fused, reference, rtol=0, atol=0) + del reference, fused + row = {"n": n, "k": k, "bf16_exact": True, + "torch": measure(dequantize_serialized_nvfp4, packed, scales, global_scale), + "fused": measure(dequantize_nvfp4_cuda, packed, scales, global_scale)} + row["speedup"] = row["torch"]["gpu_ms"] / row["fused"]["gpu_ms"] + print(json.dumps(row), flush=True) + rows.append(row) + del packed, scales + args.output.write_text(json.dumps({"gpu": torch.cuda.get_device_name(), "torch": torch.__version__, + "source_commit": os.environ.get("FASTVIDEO_SOURCE_COMMIT"), + "global_scale": global_scale, "rows": rows}, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py new file mode 100644 index 0000000000..4882b9995d --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -0,0 +1,215 @@ +"""Pruned FastH3 8-step benchmark on a single RTX 4090. + +usage: python bench_pod.py [--timed N] [--offload-buffers] [--no-layerwise] +Stage logs (FASTVIDEO_STAGE_LOGGING=1) carry per-stage time and memory peaks; results.json lands in outputs//. +""" +import argparse +import json +import os +import pathlib +import shlex +import statistics +import subprocess +import sys +import threading +import time + + +class HostMemoryPeak: + """Sample pod-wide cgroup usage; anon excludes cached checkpoint file pages.""" + + def __init__(self, cgroup_root=pathlib.Path("/sys/fs/cgroup")): + self.cgroup_root = cgroup_root + self.stop = threading.Event() + self.peak_bytes = 0 + self.peak_anon_bytes = 0 + self.peak_gpu_bytes = None + self.host_error: str | None = None + self._gpu_used = None + self._nvml_shutdown = None + self.thread = threading.Thread(target=self._sample, daemon=True) + try: + import pynvml + pynvml.nvmlInit() + self._nvml_shutdown = pynvml.nvmlShutdown + handle = pynvml.nvmlDeviceGetHandleByIndex(0) + self._gpu_used = lambda: pynvml.nvmlDeviceGetMemoryInfo(handle).used + except Exception as exc: + print(f"GPU memory sampler unavailable: {exc}", flush=True) + + def _sample(self): + while not self.stop.is_set(): + if self.host_error is None: + try: + root = self.cgroup_root + if (root / "memory.current").is_file(): + used = int((root / "memory.current").read_text()) + stats = dict(line.split() for line in (root / "memory.stat").read_text().splitlines()) + anon = int(stats["anon"]) + else: + root = root / "memory" + used = int((root / "memory.usage_in_bytes").read_text()) + stats = dict(line.split() for line in (root / "memory.stat").read_text().splitlines()) + # cgroup v1 RSS counts anonymous memory, excluding file cache. + anon = int(stats["total_rss"] if "total_rss" in stats else stats["rss"]) + self.peak_bytes = max(self.peak_bytes, used) + self.peak_anon_bytes = max(self.peak_anon_bytes, anon) + except (OSError, KeyError, ValueError) as exc: + self.host_error = f"{type(exc).__name__}: {exc}" + print(f"host memory sampling stopped: {self.host_error}", flush=True) + # Capacity verification must continue even if host counters are unavailable. + if self._gpu_used is not None: + self.peak_gpu_bytes = max(self.peak_gpu_bytes or 0, self._gpu_used()) + self.stop.wait(0.1) + + def __enter__(self): + self.thread.start() + return self + + def __exit__(self, *_args): + self.stop.set() + self.thread.join() + if self._nvml_shutdown is not None: + self._nvml_shutdown() + + def metrics(self): + return { + "peak_gpu_used_gib": (round(self.peak_gpu_bytes / 2**30, 3) + if self.peak_gpu_bytes is not None else None), + "peak_host_cgroup_gib": (round(self.peak_bytes / 2**30, 3) if self.host_error is None else None), + "peak_host_anon_gib": (round(self.peak_anon_bytes / 2**30, 3) if self.host_error is None else None), + "host_memory_error": self.host_error, + } + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("name"); ap.add_argument("model"); ap.add_argument("quant", choices=("nvfp4", "fp8", "bf16")) + ap.add_argument("--timed", type=int, default=2) + ap.add_argument("--offload-buffers", action="store_true") + ap.add_argument("--no-layerwise", action="store_true") + ap.add_argument("--resident-encoder", action="store_true") + ap.add_argument("--pageable-host", action="store_true", help="retain file-backed CPU offload weights without pinning") + ap.add_argument("--tile-batch", default=None) + ap.add_argument("--frames", type=int, default=243) + ap.add_argument("--seed", type=int, default=20260929) + ap.add_argument("--height", type=int, default=768) + ap.add_argument("--width", type=int, default=1344) + ap.add_argument("--warmup", type=int, default=1) + ap.add_argument("--prompt-file", type=pathlib.Path, default=pathlib.Path(__file__).with_name("prompts_1k.json")) + ap.add_argument("--output-root", type=pathlib.Path, default=pathlib.Path("/workspace/outputs")) + ap.add_argument("--no-vae-compile", action="store_true") + ap.add_argument("--adaln-cache", action="store_true") + ap.add_argument("--profile", action="store_true") + ap.add_argument("--sparsity", type=float, default=0.8) + ap.add_argument("--decode", default="h3-vae") + ap.add_argument("--lazy", action="store_true", help="lazy_module_load: reload released modules per request") + ap.add_argument("--prompts", default=None, help="comma-separated prompt ids (default: both)") + ap.add_argument("--once", action="store_true", help="render each selected prompt once, without extra warmup clips") + a = ap.parse_args() + if not a.once and (a.timed < 2 or a.warmup < 1): + ap.error("Use at least one warmup and two timed runs") + if not pathlib.Path(a.model, "fastvideo_inference.json").is_file(): + ap.error("The model directory must contain fastvideo_inference.json for the 8-step DMD contract") + + os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1") + # NVFP4 is retained for other GPUs; the RTX 4090 DiT uses FP8. + if a.quant == "nvfp4": + os.environ.setdefault("FASTVIDEO_H3_VSA_FP4", "1") + os.environ.setdefault("FASTVIDEO_NVFP4_MM_BACKEND", "cutlass") + os.environ.setdefault("FASTVIDEO_MINIMAX_H3_FUSIONS", "all") + os.environ.setdefault("FASTVIDEO_H3_VAE_TILE_BATCH", "28") + os.environ.setdefault("FASTVIDEO_VSA_TRITON", "1") + os.environ.setdefault("FASTVIDEO_VSA_SM100A", "0") + os.environ.setdefault("FASTVIDEO_FA4", "0") + os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") + ap_table = os.environ.get("FASTVIDEO_H3_ADALN_TABLE") + if a.profile: + os.environ["FASTVIDEO_H3_SP_PROFILE"] = "1" + if a.adaln_cache and not ap_table: + os.environ["FASTVIDEO_H3_ADALN_CACHE"] = "1" + # Export the exact modulation tables (and inputs) for table-only loads and low-rank experiments. + os.environ.setdefault("FASTVIDEO_H3_ADALN_DUMP", f"/workspace/adaln_tables_{a.name}.pt") + if a.offload_buffers: + os.environ["FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS"] = "1" + if a.tile_batch: + os.environ["FASTVIDEO_H3_VAE_TILE_BATCH"] = a.tile_batch + import torch + from fastvideo import VideoGenerator + + texts = json.loads(a.prompt_file.read_text()) + layerwise = not a.no_layerwise + engine = {"num_gpus": 1, "use_fsdp_inference": False, + "parallelism": {"tp_size": 1, "sp_size": 1}, + "offload": {"dit": False, "dit_layerwise": layerwise, "text_encoder": not a.resident_encoder, + "vae": layerwise, "pin_cpu_memory": not a.pageable_host, "lazy_module_load": a.lazy}, + "compile": {"enabled": False, "vae_enabled": not a.no_vae_compile}} + if a.quant == "nvfp4": + engine["quantization"] = {"transformer_quant": "NVFP4", "layer_profile": "h3_dit_vsa"} + elif a.quant == "fp8": + engine["quantization"] = {"transformer_quant": "FP8"} + experimental = {"attention_backend": "VIDEO_SPARSE_ATTN_H3", "VSA_sparsity": a.sparsity, "VSA_tile_size": 64, + "h3_sequential_load": not a.resident_encoder, "inference_torch_compile": False, + "vae_parallel_decode": False, "video_decode_backend": a.decode} + config = {"model_path": a.model, "engine": engine, "pipeline": {"experimental": experimental}} + out_dir = a.output_root / a.name + out_dir.mkdir(parents=True, exist_ok=True) + model_root = pathlib.Path(a.model) + revision_file = model_root / ".cache/huggingface/download/fastvideo_inference.json.metadata" + model_revision = revision_file.read_text().splitlines()[0] if revision_file.is_file() else None + hardware = subprocess.check_output( + ["nvidia-smi", "--query-gpu=name,memory.total,driver_version,pci.bus_id", "--format=csv,noheader"], text=True + ).strip() + sampling = {"seed": a.seed, "height": a.height, "width": a.width, "num_frames": a.frames, "fps": 24, + "num_inference_steps": 9, "guidance_scale": 1.0, "batch_cfg": False} + results = {"name": a.name, "quant": a.quant, "command": shlex.join([sys.executable, "-P", *sys.argv]), + "env": {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_")) + or k in ("CUDA_VISIBLE_DEVICES", "MAX_JOBS")}, + "torch": torch.__version__, "cuda": torch.version.cuda, + "hardware": hardware, "model_revision": model_revision, + "model_contract": json.loads((model_root / "fastvideo_inference.json").read_text()), + "source_commit": os.environ.get("FASTVIDEO_SOURCE_COMMIT"), + "gpu": torch.cuda.get_device_name(0), "config": config, "sampling": sampling, + "mode": "showcase_once" if a.once else "benchmark", "runs": []} + (out_dir / "results.json").write_text(json.dumps(results, indent=2)) + t0 = time.perf_counter() + generator = VideoGenerator.from_config(config) + results["load_s"] = round(time.perf_counter() - t0, 1) + ids = a.prompts.split(",") if a.prompts else list(texts) + order = ids if a.once else [ids[i % len(ids)] for i in range(a.warmup + a.timed)] + warmup = 1 if a.once else a.warmup + try: + for i, pid in enumerate(order): + request = {"prompt": texts[pid], "negative_prompt": "", + "sampling": sampling, + "output": {"output_path": str(out_dir / f"{i:02d}_{pid}.mp4"), "save_video": True, + "return_frames": False}} + t = time.perf_counter() + try: + with HostMemoryPeak() as host_peak: + generator.generate(request) + except Exception as exc: + failed = {"prompt": pid, "warmup": i < warmup, + "wall_s": round(time.perf_counter() - t, 2), + "error": f"{type(exc).__name__}: {exc}", **host_peak.metrics()} + results.setdefault("failed_runs", []).append(failed) + print("RUN_FAILED", json.dumps(failed), flush=True) + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + raise + wall = round(time.perf_counter() - t, 2) + results["runs"].append({"prompt": pid, "warmup": i < warmup, "wall_s": wall, + "clip": request["output"]["output_path"], + **host_peak.metrics()}) + timed = [run["wall_s"] for run in results["runs"] if not run["warmup"]] + if timed: + results["median_e2e_s"] = statistics.median(timed) + results["mean_e2e_s"] = statistics.mean(timed) + print("RUN", json.dumps(results["runs"][-1]), flush=True) + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + finally: + generator.shutdown() + print("DONE", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py new file mode 100644 index 0000000000..b03b6f70c5 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py @@ -0,0 +1,62 @@ +"""Validate and time candidate kernels on captured real H3 Q/K/V. + +Run on an idle sm89 GPU, separately from clip benchmarks. Captures come from +FASTVIDEO_H3_CAPTURE_QKV on the tile-first path; they retain two full heads. +""" +import argparse +import json +import pathlib + +import torch +import triton + +from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention +from fastvideo_kernel.block_sparse_attn import block_sparse_attn + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("captures", type=pathlib.Path) + parser.add_argument("--output", type=pathlib.Path, required=True) + args = parser.parse_args() + records = [] + with torch.inference_mode(): + for capture in sorted(args.captures.glob("layer-*.pt")): + state = torch.load(capture, map_location="cuda", weights_only=True) + q, k, v, mask, vbs = (state[key] for key in ("q", "k", "v", "mask", "vbs")) + def baseline(): + return block_sparse_attn(q, k, v, mask, vbs)[0] + expected = baseline() + for int8_qk, fp8_pv, v_tiles, dynamic_p in ((True, True, False, False), + (True, True, True, False), + (True, True, True, True), + (True, False, False, False), + (False, True, True, True), + (False, False, False, False)): + def candidate(): + return sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=int8_qk, fp8_pv=fp8_pv, fp8_v_tiles=v_tiles, fp8_dynamic_p=dynamic_p) + output = candidate() + valid_rows = state["untile"] + ref = expected.index_select(2, valid_rows).float() + actual = output.index_select(2, valid_rows).float() + delta = actual - ref + reference_ms = triton.testing.do_bench(baseline) + candidate_ms = triton.testing.do_bench(candidate) + record = {"capture": capture.name, "shape": list(q.shape), "int8_qk": int8_qk, "fp8_pv": fp8_pv, "fp8_v_tiles": v_tiles, "fp8_dynamic_p": dynamic_p, + "mask_density": float(mask.float().mean()), + "finite": bool(torch.isfinite(actual).all()), + "relative_l2": float(delta.norm() / ref.norm()), + "max_abs": float(delta.abs().max()), + "cosine": float(torch.nn.functional.cosine_similarity(actual.flatten(), ref.flatten(), dim=0)), + "bf16_ms": reference_ms, "int8_fp8_ms": candidate_ms, + "speedup": reference_ms / candidate_ms} + print(json.dumps(record), flush=True) + records.append(record) + if not records: + raise RuntimeError("No real Q/K/V captures found") + args.output.write_text(json.dumps({"gpu": torch.cuda.get_device_name(), "torch": torch.__version__, + "triton": triton.__version__, "records": records}, indent=2) + "\n") + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_4090/download_model.py b/scripts/benchmarks/minimax_h3_4090/download_model.py new file mode 100644 index 0000000000..d76718a028 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/download_model.py @@ -0,0 +1,14 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Download the private FP8 model without exposing its credential.""" + +import os +from pathlib import Path + +from huggingface_hub import snapshot_download + +os.environ.setdefault("HF_XET_HIGH_PERFORMANCE", "1") +snapshot_download( + "FastVideo/FastH3-Pruned-8Step-FP8-ckpt300", + local_dir="/workspace/vol/pruned_fp8_300", + token=Path("/root/.hf-fastvideo/token").read_text().strip(), +) diff --git a/scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py b/scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py new file mode 100644 index 0000000000..af90f211a4 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py @@ -0,0 +1,102 @@ +"""Exact V2 T2AV modulation tables for the checkpoint's fixed DMD ladder.""" +import argparse +import hashlib +import json +from pathlib import Path +from types import SimpleNamespace + +import torch +import torch.nn.functional as F +from safetensors import safe_open +from fastvideo.layers.mlp import MLP +from fastvideo.layers.visual_embedding import Timesteps +from fastvideo.models.dits.minimax_h3 import MiniMaxH3AdaLayerNormModulation +from fastvideo.models.schedulers.scheduling_minimax_h3 import MiniMaxH3Scheduler +from fastvideo.pipelines.basic.minimax_h3.packing import build_row_timesteps +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_denoising import MiniMaxH3DenoisingStage + + +@torch.inference_mode() +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('model', type=Path) + parser.add_argument('output', type=Path) + parser.add_argument('--source-commit', help='FastVideo revision used to build the table') + args = parser.parse_args() + config = json.loads((args.model / 'transformer/config.json').read_text()) + contract = json.loads((args.model / 'fastvideo_inference.json').read_text()) + if config.get('adaln_rank') is not None or contract.get('task') != 't2av': + raise ValueError('This tool supports full-rank V2 T2AV only') + files = list((args.model / 'transformer').glob('*.safetensors')) + locations = {} + for path in files: + with safe_open(path, framework='pt', device='cpu') as reader: + for key in reader.keys(): + if 'adaln_proj.linear.' in key or key.startswith('time_embedder.'): + locations[key] = path + + def tensor(key, dtype=None): + with safe_open(locations[key], framework='pt', device='cpu') as reader: + return reader.get_tensor(key).to(device='cuda', dtype=dtype) + + time_proj = Timesteps(config['freq_dim'], flip_sin_to_cos=True, downscale_freq_shift=0) + embedder = MLP(config['freq_dim'], config['time_embed_hidden_dim'], config['time_embed_dim'], + act_type='silu', dtype=torch.float32).cuda().eval() + for destination, source in (('fc_in', 'linear_1'), ('fc_out', 'linear_2')): + layer = getattr(embedder, destination) + layer.weight.data = tensor(f'time_embedder.{source}.weight', torch.float32) + layer.bias.data = tensor(f'time_embedder.{source}.bias', torch.float32) + video = MiniMaxH3Scheduler(shift=contract['video_scheduler_shift']) + audio = MiniMaxH3Scheduler(shift=contract['audio_scheduler_shift']) + stage = MiniMaxH3DenoisingStage(None, video, audio) + stage._set_dmd_schedule(contract['dmd_denoising_steps'], contract['num_inference_steps'], torch.device('cuda')) + # No reference rows; text rows share the video timestep in the real packer. + layout = SimpleNamespace(sequence_length=3, video_indices=torch.tensor([0]), + audio_indices=torch.tensor([1]), num_condition_video_rows=0, + num_condition_audio_rows=0) + inputs, embeddings = {}, {} + for vt, at in zip(video.timesteps, audio.timesteps, strict=True): + unique, _ = build_row_timesteps(layout, float(vt), float(at), float(vt), 1.0) + unique = unique.cuda() + temb = embedder(time_proj(unique).to(embedder.fc_in.weight.dtype)) + key = repr((tuple(unique.reshape(-1).tolist()), tuple(temb.shape), str(temb.dtype))) + embeddings[key] = temb + inputs[key] = F.silu(temb) + if len(inputs) != contract['transformer_forwards']: + raise ValueError('The contract must supply one unique timestep key per transformer forward') + tables = {} + for block in range(config['num_layers']): + prefix = f'transformer_blocks.{block}.adaln_proj.linear' + weight, bias = tensor(prefix + '.weight'), tensor(prefix + '.bias') + # Use the original release module as an independent arithmetic check. + with torch.device('meta'): + reference = MiniMaxH3AdaLayerNormModulation(config['time_embed_dim'], config['hidden_size']) + reference.linear.weight = torch.nn.Parameter(weight, requires_grad=False) + reference.linear.bias = torch.nn.Parameter(bias, requires_grad=False) + tables[block] = {} + for key, x in inputs.items(): + result = F.linear(x.to(weight.dtype), weight, bias) + expected = torch.cat(reference(embeddings[key]), dim=-1).view_as(result) + torch.testing.assert_close(result, expected, atol=0, rtol=0) + tables[block][key] = result.cpu() + del reference, weight, bias + print('BLOCK_VALIDATED', block, flush=True) + args.output.parent.mkdir(parents=True, exist_ok=True) + torch.save(tables, args.output) + torch.save({key: value.cpu() for key, value in inputs.items()}, str(args.output) + '.inputs') + revision_file = args.model / '.cache/huggingface/download/fastvideo_inference.json.metadata' + metadata = {'config_sha256': hashlib.sha256((args.model / 'transformer/config.json').read_bytes()).hexdigest(), + 'contract_sha256': hashlib.sha256((args.model / 'fastvideo_inference.json').read_bytes()).hexdigest(), + 'model_revision': revision_file.read_text().splitlines()[0] if revision_file.is_file() else None, + 'source_commit': args.source_commit, 'torch': str(torch.__version__), 'cuda': torch.version.cuda, + 'gpu': torch.cuda.get_device_name(0), + 'helper_sha256': hashlib.sha256(Path(__file__).read_bytes()).hexdigest(), + 'contract': contract, 'blocks': len(tables), 'timestep_keys': list(inputs), + 'table_sha256': hashlib.sha256(args.output.read_bytes()).hexdigest(), + 'validation': 'Every block/rung exactly equals the original release modulation module on sm89.'} + Path(str(args.output) + '.json').write_text(json.dumps(metadata, indent=2) + '\n') + print('TABLE_DONE', args.output, flush=True) + + +if __name__ == '__main__': + main() diff --git a/scripts/benchmarks/minimax_h3_4090/prompts_1k.json b/scripts/benchmarks/minimax_h3_4090/prompts_1k.json new file mode 100644 index 0000000000..553a5a7d67 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/prompts_1k.json @@ -0,0 +1,4 @@ +{ + "latency-ceramics-005": "In a quiet pottery studio an adult potter steadies a small spinning bowl while an adult apprentice watches. The apprentice asks, \"Is the rim ready?\" The potter says, \"One more gentle pass,\" and smooths the lip with a damp sponge. Begin with a close view of the hands, then make one restrained cut to a shoulder-level view showing both faces. The wheel hum, damp clay, a light splash and breathing form the soundscape. The movement is careful and unhurried, with no background music. The entire event is one finishing pass on the same bowl, not a demonstration of the whole pottery process.\nThe apprentice's apron pocket is indigo, with a plain surface. This detail belongs to the existing object, stays at its established location and remains subordinate to the main action. Preserve its material and appearance through the camera movement.\nThe wheel sits near a tall north-facing window on the left wall of the studio.\nIts circular splash pan is at the seated potter's waist.\nThe apprentice stands beyond the right side of the wheel, with enough room between their clothes and the wet rim to avoid accidental contact.\nBehind the potter, a shallow shelf holds a few dry cups.\nA sink is farther back on the right, beneath a small rack of towels.\nThese positions remain fixed through the cut.\nThe wider view reveals the same work area that surrounds the hands in the opening close view.\nThe bowl is modest in size, comfortably held between two hands.\nIts side rises from a low foot in a continuous shallow curve and ends in a slightly thickened lip.\nThe clay is warm gray with a brown undertone, darker where it is wet.\nA narrow spiral line below the rim records the potter's earlier touch.\nThat line rotates with the bowl, while a stationary highlight from the window travels over the passing surface.\nKeep these two motions distinguishable: the material marks belong to the clay, and the reflected light belongs to the relationship between the surface and the window.\nOpen with the camera low enough to see the bowl's interior without looking directly down onto the wheel.\nThe potter's left fingertips support the inside edge.\nThe right hand holds a small natural sponge outside the lip.\nNeither hand blocks the whole form.\nThe near edge of the splash pan occupies the bottom of the composition as a soft curved boundary.\nBeyond the hands, the apprentice's apron is initially out of focus.\nThis arrangement gives the close shot depth and prepares the later view of the two people without requiring another establishing shot.\nThe left hand is already stable when the clip begins.\nIts fingers form a loose supportive curve rather than squeezing the wall.\nThe sponge approaches the outside edge with only a small adjustment of the wrist.\nAs it contacts the clay, it compresses slightly and darkens where moisture gathers.\nThe bowl continues to rotate at a steady moderate speed.\nShow the finishing pass as a change in the surface's smoothness and the evenness of the lip, not as a large change in the bowl's overall shape.\nThe work is nearly finished before this moment begins.\nThe apprentice asks the question while looking at the rim, then briefly lifts their eyes toward the potter.\nTheir hands rest loosely together in front of the apron, safely away from the rotating work.\nThe question is curious and quiet, with the natural upward inflection of someone checking a detail.\nThe potter answers without stopping the wheel or turning their whole body.\nA small glance toward the apprentice is sufficient before attention returns to the clay.\nKeep the spoken words exactly as given, with no narrator explaining the technique and no extra exchange after the answer.\nCut once after the question has made the apprentice's presence clear.\nThe shoulder-level view places the potter to the left and the apprentice to the right, preserving the established relation to the wheel.\nThe bowl remains visible between them in the lower part of the frame.\nThe potter's right hand still holds the same sponge at the same point on the rim.\nContinue the wheel sound across the cut without a restart.\nThe change of view should feel like a closer understanding of the same instant, not a jump forward to a later stage of the work.\nThe potter wears a practical cotton work shirt with the sleeves rolled above the wrists.\nThe folds gather at the elbows and remain dry there.\nSmall clay marks on the forearms and apron are concentrated near the work area.\nThey do not spread or migrate during the pass.\nThe apprentice's apron is cleaner but shows a few dry pale smudges near one pocket.\nBoth garments have weight and ordinary creases.\nAvoid pristine costumes or exaggerated distressing; this is a used studio where people work carefully and clean their tools regularly.\nGive the potter a focused, patient expression.\nTheir mouth moves only for the brief reply, then settles while they feel the rim through the sponge.\nAllow the last small movement to settle within the established composition, with the environmental sound continuing around it. Preserve the quiet final composition.", + "latency-harbor-005": "## Harbor: the tide chart\nThe fabric cover around the folded tide chart is indigo, with a plain surface. This detail belongs to the existing object, stays at its established location and remains subordinate to the main action. Preserve its material and appearance through the camera movement.\nThe event takes place beside a small passenger ferry tied to a working harbor pier just before sunrise. An adult mechanic stands on the ferry's open side deck, and its captain stands beside the cabin entrance. The mechanic offers a folded tide chart and says, \"The channel is clear.\" The captain accepts it, answers, \"Then we can go,\" and looks out toward the harbor entrance. A slow lateral camera move reveals the channel beyond their shoulders. The boat remains moored throughout this brief exchange. Close voices, water against the hull, a loose halyard and a distant gull make the soundscape; there is no music.\nThe ferry is a practical coastal launch with a dark blue hull and a narrow cream band beneath its windows.\nIt carries a small enclosed cabin forward and an open passenger area behind it.\nThe camera is on the open deck, looking diagonally toward the cabin and the gap between the two people.\nThis angle places the pier along the left edge of the view and open water farther to the right.\nThe horizon is low enough that the upper part of the cabin has a clear silhouette against the pale sky.\nNothing in the composition suggests that the ferry is already underway.\nThe mechanic has finished a routine inspection rather than an emergency repair.\nTheir expression is alert but comfortable, with the slight tiredness of an early start.\nThey wear a plain work jacket over a warm shirt and carry no conspicuous badge or brand.\nA few old creases in the jacket show where the elbows bend.\nThe sleeve nearest the chart has a darker damp patch near its cuff from resting against the rail.\nKeep that patch in the same place as the arm moves.\nThe mechanic's free hand rests lightly on the top of a closed tool bag at hip level.\nThe captain is a different adult, dressed for a cool morning outside.\nA heavy knit sweater is visible beneath an open weatherproof coat.\nTheir hair is tidy but not freshly styled, and the light catches a few loose strands when they turn toward the water.\nTheir stance is balanced on the gently moving deck, with one foot slightly ahead of the other.\nThey are listening to the mechanic before the first line begins.\nThe captain does not interrupt or make a broad theatrical gesture.\nTheir reply is a small decision shared between people accustomed to working together.\nThe tide chart is a real paper object with several old folds.\nIt is partly folded into a rectangle that can be held in one hand, but one narrow flap remains loose.\nFaint printed lines and numbers are visible as a texture on its pale surface without becoming a readable title or a map inset.\nA soft graphite mark near one fold suggests that it has been used for planning.\nThe mechanic holds its lower edge between the thumb and fingers, keeping the paper clear of the damp rail.\nIts upper corner lifts slightly in the breeze before the captain takes it.\nBegin with both people already in the frame.\nThe mechanic's hand and the chart occupy the space between their bodies, below their faces.\nThis arrangement lets the first line and the handover belong to the same view.\nAs the mechanic speaks, the chart moves a short distance toward the captain.\nThe motion is neither a flourish nor an abrupt thrust.\nThe captain's receiving hand rises from beside the coat, touches the opposite edge and supports it before the mechanic releases their grip.\nThe paper bends a little between the two hands during that shared moment of support.\nThe mechanic's line is spoken in an ordinary low voice suitable for the quiet morning. The consonants remain clear, but the delivery does not sound like a public announcement. Their mouth and jaw make the small movements of the exact words, and their eyes remain on the captain. There is a slight release of breath after \"clear.\" The captain acknowledges the information with a very small nod before replying. The pause is long enough to register listening and short enough that the exchange feels familiar. Do not add another question, greeting or explanation of the voyage.\nWhen the captain says, \"Then we can go,\" the first part of the line is addressed to the mechanic. On the last words, their gaze begins to move toward the channel. The head follows the eyes through a modest turn, revealing more of the cheek nearest the exterior light. The chart settles against the front of the coat, still visibly held. The mechanic follows the captain's look with a quieter change of attention. Both remain in place. The ending is anticipation of departure, not departure itself: no engine surge, released rope or sudden movement of the ferry is needed.\nThe camera makes a restrained lateral movement toward the open-water side of the deck.\nKeep the final gesture restrained and preserve the surrounding atmosphere. Preserve the quiet final composition as the scene reaches its stated resolution." +} \ No newline at end of file diff --git a/scripts/benchmarks/minimax_h3_4090/summarize.py b/scripts/benchmarks/minimax_h3_4090/summarize.py new file mode 100644 index 0000000000..fffecde50e --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/summarize.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Join benchmark results with per-request stage logs, excluding warmup runs.""" + +import argparse +import json +import re +import statistics +from pathlib import Path + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("results", type=Path) + parser.add_argument("log", type=Path) + args = parser.parse_args() + results = json.loads(args.results.read_text()) + completed = [] + stages = {} + peaks = [] + for line in args.log.read_text().splitlines(): + timing = re.search(r"\[(\w+)_stage\|[^\]]+\] Execution completed in ([\d.]+) ms", line) + if timing: + stages[timing[1]] = float(timing[2]) / 1000 + peak = re.search(r"Memory peak_allocated=([\d.]+) GiB", line) + if peak: + peaks.append(float(peak[1])) + if line.startswith("RUN "): + run = json.loads(line[4:]) + run["stage_s"] = stages + run["peak_gpu_allocated_gib"] = max(peaks) if peaks else None + completed.append(run) + stages = {} + peaks = [] + if len(completed) != len(results["runs"]): + raise ValueError("Log and results.json have different completed run counts") + timed = [run for run in completed if not run["warmup"]] + if len(timed) < 2: + raise ValueError("At least two timed runs are required for a baseline summary") + for run in timed: + if not {"denoising", "video_decoding", "audio_decoding"}.issubset(run["stage_s"]): + raise ValueError("Missing stage timings for a completed run") + summary = { + "name": results["name"], + "sampling": results["sampling"], + "source_commit": results["source_commit"], + "timed_runs": len(timed), + "median_e2e_s": statistics.median(run["wall_s"] for run in timed), + "median_stage_s": {name: statistics.median(run["stage_s"][name] for run in timed) + for name in ("conditioning", "denoising", "video_decoding", "audio_decoding")}, + "peak_gpu_allocated_gib": max(run["peak_gpu_allocated_gib"] for run in timed), + "peak_gpu_used_gib": max((run["peak_gpu_used_gib"] for run in timed + if run.get("peak_gpu_used_gib") is not None), default=None), + "peak_host_anon_gib": max(run["peak_host_anon_gib"] for run in timed), + "peak_host_cgroup_gib": max(run["peak_host_cgroup_gib"] for run in timed), + "notes": "Stage times include deferred component loading. Host peaks are pod-wide samples every 100 ms.", + "runs": timed, + } + args.results.with_name(f"{args.results.stem}-summary.json").write_text(json.dumps(summary, indent=2) + "\n") + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_pro6000/a2a8.py b/scripts/benchmarks/minimax_h3_pro6000/a2a8.py new file mode 100644 index 0000000000..f2659517eb --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/a2a8.py @@ -0,0 +1,72 @@ +"""All-to-all bandwidth on 8x RTX PRO 6000 at FastH3 Ulysses SP8 payload sizes (10 s, 1344x768: 73.6k tokens).""" +import modal + +image = (modal.Image.debian_slim(python_version="3.12") + .pip_install("torch==2.8.0", index_url="https://download.pytorch.org/whl/cu128")) +app = modal.App("h3-a2a8", image=image) + +TOKENS, WIDTH = 73642, 7168 +WORLD = 8 + + +def _worker(rank: int, env: dict, out_q) -> None: + import os + import time + import torch + import torch.distributed as dist + os.environ.update(env) + os.environ.update({"MASTER_ADDR": "127.0.0.1", "MASTER_PORT": "29511"}) + torch.cuda.set_device(rank) + dist.init_process_group("nccl", rank=rank, world_size=WORLD) + rows = -(-TOKENS // WORLD) + payloads = { + "qkvg_bf16": 4 * rows * WIDTH * 2, # q, k, v, gate shards before attention + "qkv_fp4_plus_gate_bf16": 3 * rows * WIDTH * 9 // 16 + rows * WIDTH * 2, + "qkvg_fp4": 4 * rows * WIDTH * 9 // 16, # 4-bit values + one e4m3 scale per 16 + "out_bf16": rows * WIDTH * 2, # attention output back to sequence shards + } + res = {} + for name, nbytes in payloads.items(): + n = (nbytes // 2 // WORLD) * WORLD + send = torch.empty(n, dtype=torch.bfloat16, device="cuda") + recv = torch.empty_like(send) + for _ in range(3): + dist.all_to_all_single(recv, send) + torch.cuda.synchronize() + dist.barrier() + iters = 10 + t = time.perf_counter() + for _ in range(iters): + dist.all_to_all_single(recv, send) + torch.cuda.synchronize() + ms = (time.perf_counter() - t) / iters * 1e3 + res[name] = {"mb_per_rank": round(n * 2 / 1e6, 1), "ms": round(ms, 2), + "algbw_gbps": round(n * 2 / ms / 1e6, 1)} + if rank == 0: + out_q.put(res) + dist.destroy_process_group() + + +@app.function(gpu="RTX-PRO-6000:8", cpu=16, memory=65536, timeout=1800) +def bench() -> dict: + import subprocess + import torch.multiprocessing as mp + report = {"topo": subprocess.run(["nvidia-smi", "topo", "-m"], capture_output=True, text=True).stdout[-3000:]} + for label, env in (("p2p_default", {}), ("p2p_disabled", {"NCCL_P2P_DISABLE": "1"})): + ctx = mp.get_context("spawn") + q = ctx.Queue() + procs = [ctx.Process(target=_worker, args=(r, env, q)) for r in range(WORLD)] + for p in procs: + p.start() + report[label] = q.get(timeout=900) + for p in procs: + p.join() + return report + + +@app.local_entrypoint() +def main(): + import json + r = bench.remote() + print(r.pop("topo")) + print("A2A", json.dumps(r, indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py new file mode 100644 index 0000000000..3ece05c8b3 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -0,0 +1,610 @@ +"""FastH3 V2 NVFP4 on one RTX PRO 6000 (sm_120): build, convert, time generations.""" +import json +import os +import pathlib +import subprocess +import time + +import modal + +_HERE = pathlib.Path(__file__).resolve() +# Inside the container this file is /root/app.py; the repo is already baked into the image there. +WORKTREE = _HERE.parents[3] if len(_HERE.parents) > 3 else pathlib.Path("/src/fastvideo") +CUTLASS_COMMIT = "e67e63c331d6e4b729047c95cf6b92c8454cba89" +volume = modal.Volume.from_name("h3-pro6000-weights", create_if_missing=True) + +image = ( + modal.Image.from_registry("nvidia/cuda:13.0.1-devel-ubuntu24.04", add_python="3.12") + .apt_install("git", "build-essential", "ffmpeg", "libgl1", "libglib2.0-0") + .pip_install("uv") + .add_local_dir(WORKTREE, "/src/fastvideo", copy=True, + ignore=[".git", "**/__pycache__", "fastvideo-kernel/include/cutlass/**", + "fastvideo-kernel/include/tk/**", "fastvideo/third_party/eval/**", "docs/**", + "assets/**", "comfyui/**", "apps/**", "**/*.mp4", "**/*.log"]) + .run_commands("cd /src/fastvideo && UV_TORCH_BACKEND=cu130 uv pip install --system -e . --no-sources") + .run_commands("uv pip install --system 'cmake==3.31.6' ninja 'scikit-build-core>=0.10' pybind11 hf_transfer") + .run_commands(f"git clone --filter=blob:none https://github.com/NVIDIA/cutlass.git /cutlass && " + f"git -C /cutlass checkout {CUTLASS_COMMIT}") + .env({"FLASHINFER_CUDA_ARCH_LIST": "12.0a", "FLASHINFER_WORKSPACE_BASE": "/vol/cache/flashinfer", + "TORCHINDUCTOR_CACHE_DIR": "/vol/cache/inductor", "TRITON_CACHE_DIR": "/vol/cache/triton", + "FASTVIDEO_VSA_SM100A": "0", "FASTVIDEO_FA4": "0", "FASTVIDEO_STAGE_LOGGING": "1", + "HF_HUB_ENABLE_HF_TRANSFER": "1"}) +) +app = modal.App("h3-pro6000-fastest", image=image) + + +def _sh(cmd: str, **kw) -> str: + proc = subprocess.run(cmd, shell=True, capture_output=True, text=True, **kw) + out = (proc.stdout + proc.stderr)[-8000:] + if proc.returncode != 0: + raise RuntimeError(f"command failed ({proc.returncode}): {cmd}\n{out}") + return out + + +@app.function(cpu=16, memory=32768, timeout=3600, volumes={"/vol": volume}) +def build_kernel() -> str: + _sh("rm -rf /src/fastvideo/fastvideo-kernel/include/cutlass && " + "ln -s /cutlass /src/fastvideo/fastvideo-kernel/include/cutlass && mkdir -p /vol/wheels/cu130") + _sh("rm -f /vol/wheels/cu130/*.whl") + env = dict(os.environ, TORCH_CUDA_ARCH_LIST="12.0a", MAX_JOBS="16", CC="gcc", CXX="g++", CUDAHOSTCXX="g++", + CMAKE_ARGS="-DFASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER=ON -DFASTVIDEO_KERNEL_BUILD_TK=OFF " + "-DGPU_BACKEND=CUDA -DCMAKE_CUDA_ARCHITECTURES=120a") + _sh("cd /src/fastvideo/fastvideo-kernel && pip wheel . --no-build-isolation --no-deps -w /vol/wheels/cu130", + env=env) + volume.commit() + return _sh("ls -la /vol/wheels/cu130") + + +def _install_kernel(): + _sh("pip install --no-deps --force-reinstall /vol/wheels/cu130/*.whl") + + +def _build_light_int8_vae() -> str: + """26-block LynnReal light decoder: dense fp16 decoder + official encoder, Kijai int8-convrot overlay.""" + from huggingface_hub import hf_hub_download + from safetensors import safe_open + from safetensors.torch import save_file + + target = pathlib.Path("/vol/fv/vae_light_int8") + if (target / "config.json").exists(): + return "exists" + target.mkdir(parents=True, exist_ok=True) + overlay = hf_hub_download("Kijai/MiniMax-H3-experimental", "minimax_h3_lynnreal_light_vae_int8_convrot.safetensors", + local_dir="/vol/kijai") + official = pathlib.Path("/vol/official/vae") + weight_map = json.loads((official / "diffusion_pytorch_model.safetensors.index.json").read_text())["weight_map"] + tensors = {} + for shard in sorted(set(weight_map.values())): + with safe_open(str(official / shard), framework="pt") as reader: + for key in reader.keys(): + if not key.startswith("decoder.transformer_blocks."): + tensors[key] = reader.get_tensor(key) + with safe_open("/vol/light-vae/lynnreal_light_vae_decoder_fp16.safetensors", framework="pt") as reader: + light_keys = list(reader.keys()) + for key in light_keys: + tensors[key] = reader.get_tensor(key) + blocks = {int(k.split(".")[2]) for k in tensors if k.startswith("decoder.transformer_blocks.")} + assert blocks == set(range(26)), sorted(blocks) + save_file(tensors, str(target / "diffusion_pytorch_model.safetensors")) + config = json.loads((official / "config.json").read_text()) + config["decoder_num_layers"] = 26 + (target / "config.json").write_text(json.dumps(config, indent=2) + "\n") + (target / "minimax_h3_video_vae_int8_convrot.safetensors").symlink_to(overlay) + return f"light decoder keys={len(light_keys)} total={len(tensors)} blocks={len(blocks)}" + + +@app.function(gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=10800, volumes={"/vol": volume}) +def convert(minimal: bool = True) -> dict: + _install_kernel() + report = {} + os.makedirs("/vol/fv", exist_ok=True) + if not os.path.exists("/vol/fv/text_encoder_nvfp4/config.json"): + report["text_encoder"] = _sh( + "cd /src/fastvideo && python scripts/checkpoint_conversion/convert_minimax_h3_text_encoder_nvfp4.py " + "--src /vol/v2-nvfp4/text_encoder --dst /vol/fv/text_encoder_nvfp4")[-1500:] + volume.commit() + for name, src, flag in (("transformer_full", "v2-nvfp4", "--quantize-attention"), + ("transformer_ffn", "v2-nvfp4", ""), + ("transformer_vsa", "v2-nvfp4", "--quantize-attention --quantize-gate"), + ("v4_transformer_vsa", "v4-nvfp4", "--quantize-attention --quantize-gate")): + if not os.path.exists(f"/vol/{src}/transformer"): + continue + if minimal and name not in ("transformer_vsa", "v4_transformer_vsa"): + continue + if not os.path.exists(f"/vol/fv/{name}/config.json"): + report[name] = _sh( + "cd /src/fastvideo && python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py " + f"--src /vol/{src}/transformer --dst /vol/fv/{name} {flag}")[-1500:] + volume.commit() + # VAE folders: official dense shards, plus Comfy's int8-convrot overlay variant. + for vae_name in (() if minimal else ("vae_dense", "vae_int8")): + target = pathlib.Path(f"/vol/fv/{vae_name}") + target.mkdir(parents=True, exist_ok=True) + for item in pathlib.Path("/vol/official/vae").iterdir(): + link = target / item.name + if not link.exists(): + link.symlink_to(item) + int8 = pathlib.Path("/vol/fv/vae_int8/minimax_h3_video_vae_int8_convrot.safetensors") + if not minimal and not int8.exists(): + int8.symlink_to("/vol/comfy/vae/minimax_h3_video_vae_int8_convrot.safetensors") + report["light_vae"] = _build_light_int8_vae() + # Model folders: V2 small components + chosen transformer / text encoder / VAE. + # 4-step VSA-0.9 model: its own manifest/schedulers, shared text encoder, VAEs and audio VAE. + root = pathlib.Path("/vol/fv/v4_vsa_light") + if pathlib.Path("/vol/v4-nvfp4/fastvideo_inference.json").exists(): + root.mkdir(parents=True, exist_ok=True) + for item in pathlib.Path("/vol/v4-nvfp4").iterdir(): + if item.name in ("transformer", "text_encoder", "vae", "audio_vae") or item.name.startswith("."): + continue + if not (root / item.name).exists(): + (root / item.name).symlink_to(item) + for comp, src in (("transformer", "/vol/fv/v4_transformer_vsa"), ("text_encoder", "/vol/fv/text_encoder_nvfp4"), + ("vae", "/vol/fv/vae_light_int8"), ("audio_vae", "/vol/v2-nvfp4/audio_vae")): + if not (root / comp).exists(): + (root / comp).symlink_to(src) + model_sets = (("v2_vsa_light", "transformer_vsa", "vae_light_int8"), + ("v2_full_light", "transformer_full", "vae_light_int8"), + ("v2_full_int8", "transformer_full", "vae_int8"), + ("v2_full_dense", "transformer_full", "vae_dense"), + ("v2_ffn_int8", "transformer_ffn", "vae_int8")) + for model, transformer, vae in (model_sets[:1] if minimal else model_sets): + root = pathlib.Path(f"/vol/fv/{model}") + root.mkdir(parents=True, exist_ok=True) + for item in pathlib.Path("/vol/v2-nvfp4").iterdir(): + if item.name in ("transformer", "text_encoder", "vae") or item.name.startswith("."): + continue + link = root / item.name + if not link.exists(): + link.symlink_to(item) + for comp, src in (("transformer", transformer), ("text_encoder", "text_encoder_nvfp4"), ("vae", vae)): + link = root / comp + if not link.exists(): + link.symlink_to(f"/vol/fv/{src}") + volume.commit() + report["models"] = _sh("ls -la /vol/fv /vol/fv/v2_vsa_light") + # Are V2's VSA compression gates trained (nonzero)? + import torch + from safetensors import safe_open + index = json.loads(pathlib.Path("/vol/v2-nvfp4/transformer/diffusion_pytorch_model.safetensors.index.json").read_text()) + gate_stats = {} + for blk in (0, 25, 49): + key = f"transformer_blocks.{blk}.attn.to_gate_compress.weight" + with safe_open(f"/vol/v2-nvfp4/transformer/{index['weight_map'][key]}", framework="pt") as reader: + w = reader.get_tensor(key).float() + gate_stats[blk] = {"abs_mean": w.abs().mean().item(), "nonzero_frac": (w != 0).float().mean().item()} + report["gate_stats"] = gate_stats + return report + + +PROMPTS = { + "kitesurf": ("A kite surfer carves hard across choppy bay water while the camera dives alongside; spray hisses off " + "the board edge, the sail flaps and snaps in the wind, and gulls cry overhead."), + "chef": ("(S1) In a bright home kitchen, a chef looks straight at the camera and says [English] Fold the eggs " + "gently and taste before you salt. A pot simmers behind her with soft bubbling and no music."), +} + + +@app.function(gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=5400, volumes={"/vol": volume}) +def run_variant(name: str, model: str, profile: str, attention: str, decode: str, vae_compile: bool, + height: int = 480, width: int = 832, warmups: int = 1, num_frames: int = 124, + env: dict | None = None, prompts: tuple = ("kitesurf", "chef"), sparsity: float = 0.8, + steps: int = 9, num_gpus: int = 1, parallel_decode: bool = False, pre_runs: tuple = (), + prompt_texts: dict | None = None, offload: dict | None = None, + experimental_extra: dict | None = None) -> dict: + os.environ.update(env or {}) + texts = {**PROMPTS, **(prompt_texts or {})} + _install_kernel() + import torch + from fastvideo import VideoGenerator + + if attention == "VIDEO_SPARSE_ATTN_H3": + os.environ["FASTVIDEO_VSA_TRITON"] = "1" + experimental = {"attention_backend": attention, "h3_sequential_load": False, "inference_torch_compile": False, + "vae_parallel_decode": parallel_decode, "video_decode_backend": decode} + if attention == "VIDEO_SPARSE_ATTN_H3": + experimental.update({"VSA_sparsity": sparsity, "VSA_tile_size": 64}) + experimental.update(experimental_extra or {}) + config = { + "model_path": f"/vol/fv/{model}", + "engine": {"num_gpus": num_gpus, "use_fsdp_inference": False, + "quantization": {"transformer_quant": "NVFP4", "layer_profile": profile}, + "parallelism": {"tp_size": 1, "sp_size": num_gpus}, + "offload": {"dit": False, "dit_layerwise": False, "text_encoder": False, "vae": False, + "pin_cpu_memory": num_gpus == 1, "lazy_module_load": False, **(offload or {})}, + "compile": {"enabled": False, "vae_enabled": vae_compile}}, + "pipeline": {"experimental": experimental}, + } + t0 = time.perf_counter() + generator = VideoGenerator.from_config(config) + load_s = time.perf_counter() - t0 + out_dir = pathlib.Path(f"/vol/outputs/{name}") + out_dir.mkdir(parents=True, exist_ok=True) + results = {"name": name, "load_s": round(load_s, 1), "env": env or {}, "shape": [height, width, num_frames], + "num_gpus": num_gpus, + "runs": []} + try: + # Optional runs at other shapes first (e.g. a 480p correctness clip), same loaded model. + for j, (ph, pw, pf, pid) in enumerate(pre_runs): + request = {"prompt": texts[pid], "negative_prompt": "", + "sampling": {"seed": 20260929, "height": ph, "width": pw, "num_frames": pf, "fps": 24, + "num_inference_steps": steps, "guidance_scale": 1.0, "batch_cfg": False}, + "output": {"output_path": str(out_dir / f"pre{j:02d}_{pid}_{ph}p.mp4"), "save_video": True, + "return_frames": False}} + t = time.perf_counter() + result = generator.generate(request) + results.setdefault("pre_runs", []).append({"shape": [ph, pw, pf], "prompt": pid, + "wall_s": round(time.perf_counter() - t, 2), + "video": getattr(result, "video_path", None)}) + print("PRE_RUN", json.dumps(results["pre_runs"][-1]), flush=True) + if pre_runs: + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + volume.commit() + # Warm every distinct prompt: prompt length changes the packed sequence, and shape-specialized + # compiled kernels would otherwise recompile inside the first timed run of each prompt. + distinct = list(dict.fromkeys(prompts)) + order = [distinct[i % len(distinct)] for i in range(warmups)] + list(prompts) + for i, pid in enumerate(order): + prompt = texts[pid] + request = {"prompt": prompt, "negative_prompt": "", + "sampling": {"seed": 20260929, "height": height, "width": width, "num_frames": num_frames, "fps": 24, + "num_inference_steps": steps, "guidance_scale": 1.0, "batch_cfg": False}, + "output": {"output_path": str(out_dir / f"{i:02d}_{pid}.mp4"), "save_video": True, + "return_frames": False}} + torch.cuda.synchronize() + t = time.perf_counter() + result = generator.generate(request) + torch.cuda.synchronize() + wall = time.perf_counter() - t + results["runs"].append({"prompt": pid, "warmup": i < warmups, "wall_s": round(wall, 2), + "generation_time_s": getattr(result, "generation_time", None), + "video": getattr(result, "video_path", None)}) + results["peak_mem_gb_device"] = _sh("nvidia-smi --query-gpu=memory.used --format=csv,noheader").strip() + finally: + generator.shutdown() + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + volume.commit() + return results + + +bench_image = image.add_local_file(pathlib.Path(__file__).parent / "bench_code.py", "/root/bench_code.py") + +SHAPES = [ + # name, prefix segments (text, audio rows), video latent tokens (t, h, w) after 1x2x2 patching + ("480p_124f", [256, 414], [37, 15, 26]), + ("768p_243f", [256, 810], [72, 24, 42]), +] + + +@app.function(gpu="RTX-PRO-6000", memory=65536, cpu=8, timeout=3600, volumes={"/vol": volume}, image=bench_image) +def bench_block() -> dict: + _install_kernel() + import sys + sys.path.insert(0, "/root") + import bench_code + report = {"check": bench_code.check_tile64()} + report.update(bench_code.run(SHAPES)) + return json.loads(json.dumps(report, default=str)) + + +@app.function(gpu="RTX-PRO-6000", memory=65536, cpu=8, timeout=1800, volumes={"/vol": volume}, image=bench_image) +def density_fn() -> dict: + _install_kernel() + import sys + sys.path.insert(0, "/root") + import bench_code + return json.loads(json.dumps(bench_code.density_study(), default=str)) + + +@app.function(gpu="RTX-PRO-6000", memory=65536, cpu=8, timeout=1800, volumes={"/vol": volume}, image=bench_image) +def kcheck_fn() -> dict: + _install_kernel() + import sys + sys.path.insert(0, "/root") + import bench_code + out = {"check": bench_code.check_tile64()} + out["density"] = bench_code.density_study() + return json.loads(json.dumps(out, default=str)) + + +@app.function(gpu="RTX-PRO-6000:2", memory=229376, cpu=16, timeout=5400, volumes={"/vol": volume}) +def run_variant2(*args, **kwargs) -> dict: + return run_variant.local(*args, **kwargs) + + +@app.function(gpu="RTX-PRO-6000:8", memory=196608, cpu=16, timeout=5400, volumes={"/vol": volume}) +def run_variant8(*args, **kwargs) -> dict: + return run_variant.local(*args, **kwargs) + + +@app.function(gpu="RTX-PRO-6000:4", memory=131072, cpu=8, timeout=5400, volumes={"/vol": volume}) +def run_variant4(*args, **kwargs) -> dict: + return run_variant.local(*args, **kwargs) + + +@app.function(gpu="RTX-PRO-6000:8", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) +def bench8_pair(prompt_texts: dict, timed: tuple, warmups: int, models: tuple = ("v2", "v2_profile")) -> list: + """8-GPU 10 s 768p V2 8-step benchmark, then a stage-profiled pass, in one container (shared compile caches).""" + specs = {"v2": ("sp8_v2_8step_768p10s_1k", "v2_vsa_light", 0.8, 9, {}, timed, warmups), + # One warm + one profiled generation; CUDA-event spans log per DiT forward as H3_STAGE_MS. + "v2_profile": ("sp8_v2_8step_profile", "v2_vsa_light", 0.8, 9, {"FASTVIDEO_H3_SP_PROFILE": "1"}, + timed[:1], 1)} + out = [] + for m in models: + name, model, sparsity, steps, extra_env, prompts, warm = specs[m] + r = run_variant.local(name, model, "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", + vae_compile=True, height=768, width=1344, num_frames=243, warmups=warm, + env={**FAST_ENV, **extra_env}, prompts=prompts, sparsity=sparsity, steps=steps, + num_gpus=8, parallel_decode=True, prompt_texts=prompt_texts) + print("RESULT", json.dumps(r), flush=True) + out.append(r) + return out + + +@app.function(gpu="RTX-PRO-6000", memory=196608, cpu=16, timeout=3 * 3600, volumes={"/vol": volume}) +def memladder(prompt_texts: dict, pid: str, configs: dict | None = None) -> list: + """V2 8-step 10 s 768p on one GPU under memory placements; stage logs carry per-stage peaks. + + ``configs`` maps a name to (offload overrides, experimental overrides, extra env). The env can set + FASTVIDEO_CUDA_MEMORY_CAP_GIB to emulate a smaller card. + """ + seq = {"h3_sequential_load": True} + lw = {"text_encoder": True, "vae": True, "dit_layerwise": True} + configs = configs or { + "A_resident": ({}, {}, {}), + "B_seq_encoder_vae_offload": ({"text_encoder": True, "vae": True}, seq, {}), + "C_plus_dit_layerwise": (lw, seq, {}), + } + out = [] + for name, (offload, extra, env_extra) in configs.items(): + try: + r = run_variant.local(f"mem_{name}", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", + decode="h3-vae", vae_compile=False, height=768, width=1344, num_frames=243, + warmups=1, env={**FAST_ENV, **env_extra}, prompts=(pid,), sparsity=0.8, steps=9, + num_gpus=1, prompt_texts=prompt_texts, offload=offload, experimental_extra=extra) + except Exception as e: # keep the ladder going; record the failure + r = {"name": f"mem_{name}", "error": repr(e)[:3000]} + print("RESULT", json.dumps(r), flush=True) + out.append(r) + for k in env_extra: + os.environ.pop(k, None) + return out + + +@app.function(cpu=8, memory=32768, timeout=1800, volumes={"/vol": volume}) +def compare_videos(a: str, b: str) -> dict: + """Frame PSNR between two MP4s on the volume (same seed and prompt).""" + import imageio.v3 as iio + import numpy as np + fa = iio.imread(a, plugin="pyav").astype(np.float32) + fb = iio.imread(b, plugin="pyav").astype(np.float32) + n = min(len(fa), len(fb)) + mse = ((fa[:n] - fb[:n]) ** 2).reshape(n, -1).mean(axis=1) + psnr = 10 * np.log10(255.0**2 / np.maximum(mse, 1e-9)) + return {"frames": [len(fa), len(fb)], "psnr_mean": float(psnr.mean()), "psnr_min": float(psnr.min())} + + +FAST_ENV = {"FASTVIDEO_H3_VSA_FP4": "1", "FASTVIDEO_MINIMAX_H3_FUSIONS": "all", "FASTVIDEO_NVFP4_MM_BACKEND": "cutlass", + "FASTVIDEO_H3_VAE_TILE_BATCH": "28"} + + +@app.local_entrypoint() +def main(step: str = "all", ladder: str = "base"): + if step == "prep_personal": + print("KERNEL", build_kernel.remote()) + print("CONVERT", json.dumps(convert.remote(minimal=True), indent=1)[:6000]) + return + if step == "sp2check": + common = dict(attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=False, height=480, width=832, + num_frames=124, warmups=0, env=FAST_ENV, prompts=("kitesurf",), sparsity=0.8, steps=9) + one = run_variant.spawn("sp1_480p", "v2_vsa_light", "h3_dit_vsa", **common) + two = run_variant2.spawn("sp2_480p", "v2_vsa_light", "h3_dit_vsa", num_gpus=2, **common) + r1, r2 = one.get(), two.get() + print("RESULT", json.dumps(r1)) + print("RESULT", json.dumps(r2)) + print("COMPARE", json.dumps(compare_videos.remote(r1["runs"][0]["video"], r2["runs"][0]["video"]))) + return + if step == "sp8": + r = run_variant8.remote( + "sp8_v2_8step", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", + vae_compile=True, height=768, width=1344, num_frames=243, warmups=1, env=FAST_ENV, + prompts=("kitesurf", "chef", "kitesurf", "chef"), sparsity=0.8, steps=9, num_gpus=8, parallel_decode=True, + pre_runs=((480, 832, 124, "kitesurf"), )) + print("RESULT", json.dumps(r)) + if r.get("pre_runs"): + print("COMPARE", json.dumps(compare_videos.remote("/vol/outputs/sp1_480p/00_kitesurf.mp4", + r["pre_runs"][0]["video"]))) + return + if step == "simfp8": + # SP=1 with and without the SP exchange's FP8 rounding, same settings as sp1_480p. + common = dict(attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=False, height=480, width=832, + num_frames=124, warmups=0, prompts=("kitesurf",), sparsity=0.8, steps=9) + plain = run_variant.spawn("sp1_480p_rerun", "v2_vsa_light", "h3_dit_vsa", env=FAST_ENV, **common) + sim = run_variant.spawn("sp1_480p_simfp8", "v2_vsa_light", "h3_dit_vsa", + env={**FAST_ENV, "FASTVIDEO_H3_SIM_SP_FP8": "1"}, **common) + r_plain, r_sim = plain.get(), sim.get() + print("RESULT", json.dumps(r_plain)) + print("RESULT", json.dumps(r_sim)) + return + if step == "memreport": + r = run_variant.remote("memreport", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", + decode="h3-vae", vae_compile=False, height=480, width=832, num_frames=124, warmups=0, + env={**FAST_ENV, "FASTVIDEO_MEMORY_REPORT": "1"}, prompts=("kitesurf",), sparsity=0.8, + steps=9, offload={"text_encoder": True, "vae": True}, + experimental_extra={"h3_sequential_load": False}) + print("RESULT", json.dumps(r)) + return + if step == "memladder": + rows = [json.loads(line) for line in open(WORKTREE.parent / "UniServe-sm120fp4" / "uniserve_eval" / "workloads" + / "fast_h3" / "latency.jsonl")] + pid = "latency-ceramics-005" + text = {r["id"]: r["prompt"] for r in rows if r["id"] == pid} + configs = None + if ladder == "caps": + seq = {"h3_sequential_load": True} + lw = {"text_encoder": True, "vae": True, "dit_layerwise": True} + configs = { + "D32_C_cap32": (lw, seq, {"FASTVIDEO_CUDA_MEMORY_CAP_GIB": "32"}), + "D24_buffers_cap24": (lw, seq, {"FASTVIDEO_CUDA_MEMORY_CAP_GIB": "24", + "FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS": "1"}), + "D16_buffers_tile8_cap16": (lw, seq, {"FASTVIDEO_CUDA_MEMORY_CAP_GIB": "16", + "FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS": "1", + "FASTVIDEO_H3_VAE_TILE_BATCH": "8"}), + } + for r in memladder.remote(text, pid, configs): + print("RESULT", json.dumps(r)) + return + if step == "bench8_pair": + # UniServe's 10 s / ~1K-token latency prompts, so the numbers line up with its published protocol. + rows = [json.loads(line) for line in open(WORKTREE.parent / "UniServe-sm120fp4" / "uniserve_eval" / "workloads" + / "fast_h3" / "latency.jsonl")] + ten = {r["id"]: r["prompt"] for r in rows if r["seconds"] == 10 and r["prompt_len"] == 1000} + ids = ("latency-ceramics-005", "latency-harbor-005") + texts = {i: ten[i] for i in ids} + for r in bench8_pair.remote(texts, (ids[0], ids[1], ids[0], ids[1]), 2): + timed = sorted(x["wall_s"] for x in r["runs"] if not x["warmup"]) + print("SUMMARY", r["name"], "timed", timed, "median", timed[len(timed) // 2] if len(timed) % 2 + else (timed[len(timed) // 2 - 1] + timed[len(timed) // 2]) / 2) + return + if step == "sp4": + r = run_variant4.remote( + "sp4_v2_8step", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", + vae_compile=True, height=768, width=1344, num_frames=243, warmups=1, env=FAST_ENV, + prompts=("kitesurf", "chef", "kitesurf", "chef"), sparsity=0.8, steps=9, num_gpus=4, parallel_decode=True, + pre_runs=((480, 832, 124, "kitesurf"), )) + print("RESULT", json.dumps(r)) + if r.get("pre_runs"): + print("COMPARE", json.dumps(compare_videos.remote("/vol/outputs/sp1_480p/00_kitesurf.mp4", + r["pre_runs"][0]["video"]))) + return + if step in ("bench8_v2", "bench8_v4"): + v2 = step == "bench8_v2" + r = run_variant8.remote( + f"sp8_{'v2_8step' if v2 else 'v4_4step'}_768p10s", "v2_vsa_light" if v2 else "v4_vsa_light", "h3_dit_vsa", + attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=True, height=768, width=1344, + num_frames=243, warmups=1, env=FAST_ENV, prompts=("kitesurf", "chef", "kitesurf", "chef"), + sparsity=0.8 if v2 else 0.9, steps=9 if v2 else 5, num_gpus=8, parallel_decode=True) + print("RESULT", json.dumps(r)) + return + if step == "build_bench": + print("KERNEL", build_kernel.remote()) + step = "bench" + if step in ("prep768", "e2e768"): + if step == "prep768": + print("KERNEL", build_kernel.remote()) + print("CONVERT", json.dumps(convert.remote(), indent=1)[:6000]) + fast = {"FASTVIDEO_H3_VSA_FP4": "1", "FASTVIDEO_MINIMAX_H3_FUSIONS": "all", + "FASTVIDEO_NVFP4_MM_BACKEND": "cutlass", "FASTVIDEO_H3_VAE_TILE_BATCH": "28"} + common = dict(attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=True, height=768, width=1344, + num_frames=243, warmups=1, env=fast, prompts=("kitesurf", "chef")) + calls = { + "v4_4step_vsa90_768p10s": run_variant.spawn("v4_4step_vsa90_768p10s", "v4_vsa_light", "h3_dit_vsa", + sparsity=0.9, steps=5, **common), + "v2_8step_vsa80_768p10s": run_variant.spawn("v2_8step_vsa80_768p10s", "v2_vsa_light", "h3_dit_vsa", + sparsity=0.8, steps=9, **common), + } + for name, call in calls.items(): + try: + print("RESULT", json.dumps(call.get())) + except Exception as exc: # noqa: BLE001 + print("FAILED", name, repr(exc)[:3000]) + return + if step == "kcheck": + print("KERNEL", build_kernel.remote()) + print("KCHECK", json.dumps(kcheck_fn.remote(), indent=1)) + return + if step == "density": + print("DENSITYRESULT", json.dumps(density_fn.remote(), indent=1)) + return + if step == "bench": + print("BENCHRESULT", json.dumps(bench_block.remote(), indent=1)) + return + if step in ("all", "build"): + print("KERNEL", build_kernel.remote()) + if step in ("all", "convert"): + print("CONVERT", json.dumps(convert.remote(), indent=1)[:6000]) + if step in ("all", "run"): + variants = [ + ("full_qatinfer_light", "v2_full_light", "h3_dit", "ATTN_QAT_INFER", "h3-vae", True), + ("full_vsa_light", "v2_full_light", "h3_dit", "VIDEO_SPARSE_ATTN_H3", "h3-vae", True), + ("full_qatinfer_int8", "v2_full_int8", "h3_dit", "ATTN_QAT_INFER", "h3-vae", True), + ("full_qatinfer_taeh3", "v2_full_light", "h3_dit", "ATTN_QAT_INFER", "taeh3", False), + ("ffn_vsa_int8", "v2_ffn_int8", "h3_dit_ffn", "VIDEO_SPARSE_ATTN_H3", "h3-vae", True), + ] + # One RTX PRO 6000 per variant, all at once; failures are returned, not raised. + for v, result in zip(variants, run_variant.starmap(variants, return_exceptions=True)): + if isinstance(result, Exception): + print("FAILED", v[0], repr(result)[:3000]) + else: + print("RESULT", json.dumps(result)) + + +# --- Headline benchmark (480p 5 s, 768p 10 s) on 1 / 4 / 8 GPUs from a FastVideo HF repo. --- +# MODAL_PROFILE=aryan5v modal run --detach app.py::headline --repo FastVideo/ --profile h3_dit_ffn --gpus 1,4,8 +HERE = pathlib.Path(__file__).resolve().parent +headline_image = (image.add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") + .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) +SECRETS = [modal.Secret.from_name("hf-fastvideo")] + + +@app.function(cpu=8, memory=32768, timeout=3600, volumes={"/vol": volume}, secrets=SECRETS, image=headline_image) +def headline_fetch(repo: str) -> str: + from huggingface_hub import snapshot_download + local = f"/vol/models/{repo.split('/')[-1]}" + snapshot_download(repo, local_dir=local, token=os.environ["HF_TOKEN"], max_workers=16) + volume.commit() + return _sh(f"du -sh {local}/*") + + +def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None, tag: str = "") -> dict: + _install_kernel() + model = f"/vol/models/{repo.split('/')[-1]}" + # The run name is a single path component locally and on the volume. + run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}{tag}".replace("/", "-") + env = {**os.environ, **FAST_ENV, **(extra_env or {}), "HEADLINE_OUT": "/vol/outputs/headline", + "HEADLINE_DEVICE": f"{gpus}x RTX PRO 6000", "PYTHONPATH": "/src/fastvideo"} + proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, + "--prompts", "/root/headline_prompts.json"], env=env, capture_output=True, text=True, + cwd="/root") + log_text = proc.stdout + proc.stderr + run_dir = pathlib.Path("/vol/outputs/headline") / run_name + run_dir.mkdir(parents=True, exist_ok=True) + (run_dir / "run.log").write_text(log_text) + volume.commit() + errors = [line for line in log_text.splitlines() + if any(k in line for k in ("Error", "error:", "Traceback", "Killed", "OOM", "out of memory"))][-40:] + result_path = run_dir / "results.json" + results = json.loads(result_path.read_text()) if result_path.exists() else {} + return {"run": run_name, "returncode": proc.returncode, "results": results, "errors": errors, + "log_tail": log_text[-1500:]} + + +@app.function(image=headline_image, gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=2 * 3600, volumes={"/vol": volume}) +def headline1(repo: str, profile: str, extra_env: dict | None = None, tag: str = "") -> dict: + return _headline(repo, 1, profile, extra_env, tag) + + +@app.function(image=headline_image, gpu="RTX-PRO-6000:4", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) +def headline4(repo: str, profile: str, extra_env: dict | None = None, tag: str = "") -> dict: + return _headline(repo, 4, profile, extra_env, tag) + + +@app.function(image=headline_image, gpu="RTX-PRO-6000:8", memory=262144, cpu=32, timeout=2 * 3600, volumes={"/vol": volume}) +def headline8(repo: str, profile: str, extra_env: dict | None = None, tag: str = "") -> dict: + return _headline(repo, 8, profile, extra_env, tag) + + +@app.local_entrypoint() +def headline(repo: str, profile: str = "h3_dit_ffn", gpus: str = "1,4,8", skip_fetch: bool = False, + extra_env: str = "{}", tag: str = ""): + if not skip_fetch: + print(headline_fetch.remote(repo)) + fns = {"1": headline1, "4": headline4, "8": headline8} + calls = [fns[g].spawn(repo, profile, json.loads(extra_env), tag) for g in gpus.split(",")] + out = HERE / "headline_results" + out.mkdir(exist_ok=True) + for call in calls: + res = call.get() + print(json.dumps({k: v for k, v in res.items() if k != "log_tail"}, indent=1)[:3000]) + if res["returncode"]: + print(res["log_tail"]) + (out / f"{res['run']}.json").write_text(json.dumps(res, indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_code.py b/scripts/benchmarks/minimax_h3_pro6000/bench_code.py new file mode 100644 index 0000000000..c45557b839 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_code.py @@ -0,0 +1,338 @@ +"""Per-component timing of one MiniMax-H3 block on the current GPU. Imported inside the Modal container.""" +import json +import math +import time + +import torch +import torch.nn.functional as F + +HID, HEADS, HD, FFN = 5376, 56, 128, 14336 + + +def timed(fn, iters=5, warmup=2): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + return round(start.elapsed_time(end) / iters, 3) + + +def smooth_tokens(prefix_len, shape, dev, gen, noise=0.5): + """Token features with spatial-temporal correlation (proxy for real q/k), packed [prefix | video].""" + t, h, w = shape + low = torch.randn(1, HEADS * HD, max(2, t // 6), max(2, h // 6), max(2, w // 6), device=dev, generator=gen) + vid = F.interpolate(low, size=(t, h, w), mode="trilinear", align_corners=False) + vid = vid.reshape(HEADS, HD, -1).permute(2, 0, 1) + vid = vid + noise * torch.randn(vid.shape, device=dev, generator=gen) + pre = torch.randn(prefix_len, HEADS, HD, device=dev, generator=gen) + return torch.cat([pre, vid]).to(torch.bfloat16)[None] # [1, L, H, D] + + +def run(shapes): + from flashinfer import SfLayout, mm_fp4, nvfp4_quantize + from fastvideo.attention.backends.video_sparse_attn_h3 import (_build_block_mask, _h3_tile_geometry, _pool_tiles) + from fastvideo_kernel.block_sparse_attn import block_sparse_attn as bsa64 + from fastvideo.models.dits.minimax_h3_fusions import (fused_qknorm_rope, fused_residual_gate_rmsnorm_modulate, + fused_rmsnorm_modulate, minimax_h3_swiglu) + import attn_qat_infer.api as fa + from attn_qat_infer.api import BLOCK_M, sageattn_blackwell, sageattn_blackwell_sparse + + dev = torch.device("cuda") + gen = torch.Generator(device=dev).manual_seed(0) + report = {"device": torch.cuda.get_device_name(0), "torch": torch.__version__} + unit = torch.tensor(1.0, device=dev) + + def fp4_weight(n, k): + w = torch.randn(n, k, device=dev, dtype=torch.bfloat16) * 0.02 + gsf = (448 * 6) / w.float().abs().max() + wq, ws = nvfp4_quantize(w, gsf, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + return w, wq, ws, (1.0 / gsf).float() + + for name, prefix_segments, video_shape in shapes: + prefix_segments = tuple(prefix_segments) + video_shape = tuple(video_shape) + L = sum(prefix_segments) + math.prod(video_shape) + r = {"tokens": L} + torch.cuda.empty_cache() + + # ---------------- linears ---------------- + lin = {} + for lname, n, k in (("qkv_one", 7168, HID), ("out", HID, 7168), ("fc_in", 2 * FFN, HID), + ("fc_out", HID, FFN)): + w, wq, ws, alpha = fp4_weight(n, k) + x = torch.randn(L, k, device=dev, dtype=torch.bfloat16) + xq, xs = nvfp4_quantize(x, unit, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + entry = {"quant_ms": timed(lambda: nvfp4_quantize(x, unit, sfLayout=SfLayout.layout_128x4, + do_shuffle=False))} + for backend in ("cudnn", "cutlass", "auto"): + try: + entry[f"mm_{backend}_ms"] = timed( + lambda: mm_fp4(xq, wq.T, xs, ws.T, alpha, torch.bfloat16, None, backend=backend)) + except Exception as exc: # noqa: BLE001 + entry[f"mm_{backend}_ms"] = f"ERR {type(exc).__name__}: {str(exc)[:120]}" + entry["bf16_ms"] = timed(lambda: F.linear(x, w)) + entry["tflops_fp4_auto"] = (round(2 * L * n * k / entry["mm_auto_ms"] / 1e9, 1) + if isinstance(entry["mm_auto_ms"], float) else None) + lin[lname] = entry + del w, wq, ws, x, xq, xs + r["linear"] = lin + + # ---------------- attention ---------------- + geom = _h3_tile_geometry(prefix_segments, video_shape, dev, (4, 4, 4)) + _, vbs, untile, n_prefix, n_video = geom + n_tiles = vbs.numel() + Lpad = n_tiles * 64 + att = {"tiles64": n_tiles, "prefix_tiles": n_prefix} + prefix_len = sum(prefix_segments) + q = smooth_tokens(prefix_len, video_shape, dev, gen) + k = (q + 0.5 * torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16)).to(torch.bfloat16) + v = torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16) + g = torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16) * 0.1 + + def tile(x): + buf = torch.zeros((1, Lpad, HEADS, HD), device=dev, dtype=x.dtype) + buf[:, untile] = x + return buf + + qt, kt, vt, gt = tile(q), tile(k), tile(v), tile(g) + att["tile_scatter4_gather1_ms"] = timed(lambda: (tile(q), tile(k), tile(v), tile(g), qt[:, untile])) + + def vsa_full(sparsity=0.8, with_gate=True): + qp = _pool_tiles(qt, vbs, 64) + kp = _pool_tiles(kt, vbs, 64) + scores = torch.matmul(qp, kp.transpose(-2, -1)) / HD**0.5 + mask = _build_block_mask(scores, n_prefix, sparsity, True, ((n_prefix, n_prefix + n_video), ), (sparsity, )) + out, _ = bsa64(qt.transpose(1, 2).contiguous(), kt.transpose(1, 2).contiguous(), + vt.transpose(1, 2).contiguous(), mask, vbs) + out = out.transpose(1, 2).contiguous() + if with_gate: + vp = _pool_tiles(vt, vbs, 64) + oc = torch.matmul(torch.softmax(scores, dim=-1), vp).permute(0, 2, 1, 3).to(out.dtype) + out = (out.view(1, n_tiles, 64, HEADS, HD) + oc.unsqueeze(2) * gt.view(1, n_tiles, 64, HEADS, HD)).view( + 1, Lpad, HEADS, HD) + return out, mask + + _, mask = vsa_full() + att["tile_density"] = round(mask.float().mean().item(), 4) + att["vsa_triton_total_ms"] = timed(vsa_full) + qb, kb, vb = (x.transpose(1, 2).contiguous() for x in (qt, kt, vt)) + att["vsa_triton_kernel_only_ms"] = timed(lambda: bsa64(qb, kb, vb, mask, vbs)) + + # dense FP4 on the unpadded packed sequence + qd, kd, vd = (x.transpose(1, 2).contiguous() for x in (q, k, v)) + att["sage3_dense_fp4_ms"] = timed(lambda: sageattn_blackwell(qd, kd, vd), iters=3) + # same with delta_s cached + per_block_mean=False (no 9.5 GB memset) + zero_ds = {} + + def sage_nods(qx, kx, vx): + QL, KL = qx.size(2), kx.size(2) + qx, kx, vx = (F.pad(x, (0, 0, 0, (BLOCK_M - x.size(2) % BLOCK_M) % BLOCK_M)).contiguous() for x in (qx, kx, vx)) + key = (qx.shape[0], qx.shape[1], kx.shape[2]) + if key not in zero_ds: + zero_ds[key] = torch.zeros((qx.shape[0], qx.shape[1], 1, kx.shape[2]), device=dev, dtype=torch.float32) + ql, kl, vl = fa.scale_and_quant_fp4(qx), fa.scale_and_quant_fp4_permute(kx), fa.scale_and_quant_fp4_transpose(vx) + return fa.blockscaled_fp4_attn(ql, kl, vl, zero_ds[key], KL, False, False, True, True, None)[0][:, :, :QL] + + att["sage3_dense_fp4_no_deltas_ms"] = timed(lambda: sage_nods(qd, kd, vd), iters=3) + out_a = sageattn_blackwell(qd, kd, vd) + out_b = sage_nods(qd, kd, vd) + att["no_deltas_max_abs_diff"] = (out_a.float() - out_b.float()).abs().max().item() + Lq = qd.shape[2] + att["dense_fp4_tflops"] = round(4 * Lq * Lq * HD * HEADS / att["sage3_dense_fp4_no_deltas_ms"] / 1e9, 1) + if L < 30000: + att["sdpa_bf16_dense_ms"] = timed(lambda: F.scaled_dot_product_attention(qd, kd, vd), iters=3) + + # sparse FP4 on 128x128 blocks with 64x64 quadrant masks (exact VSA tile-64 semantics) + nt2 = n_tiles + (n_tiles % 2) + qs, ks, vs = (F.pad(x.transpose(1, 2), (0, 0, 0, nt2 * 64 - Lpad)).contiguous() for x in (qt, kt, vt)) + q2k_idx, q2k_num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, 64, vbs, validate=True) + att["block128_density"] = round(q2k_num.float().mean().item() / (nt2 // 2), 4) + att["sparse_fp4_quad_ms"] = timed(lambda: sageattn_blackwell_sparse(qs, ks, vs, q2k_idx, q2k_num, kvv, quad)) + att["mask_to_blocks_ms"] = timed(lambda: fa.vsa_tile_mask_to_fp4_blocks(mask, 64, vbs)) + o_fp4 = sageattn_blackwell_sparse(qs, ks, vs, q2k_idx, q2k_num, kvv, quad)[:, :, :Lpad] + o_tri, _ = bsa64(qb, kb, vb, mask, vbs) + rows = untile # valid (non-pad) rows of the tile buffer + diff = (o_fp4[:, :, rows].float() - o_tri[:, :, rows].float()) + att["quad_vs_triton_rel_l2"] = round((diff.norm() / o_tri[:, :, rows].float().norm()).item(), 4) + od = sageattn_blackwell(qd, kd, vd) + from torch.nn.attention import SDPBackend, sdpa_kernel + if L < 30000: + ref_dense = F.scaled_dot_product_attention(qd, kd, vd) + att["dense_fp4_vs_bf16_rel_l2"] = round(((od.float() - ref_dense.float()).norm() / + ref_dense.float().norm()).item(), 4) + att["quad_finite"] = bool(torch.isfinite(o_fp4[:, :, rows]).all()) + r["attention"] = att + del q, k, v, g, qt, kt, vt, gt, qd, kd, vd, qs, ks, vs, qb, kb, vb, out_a, out_b, zero_ds, o_fp4, o_tri, od + + # ---------------- elementwise ---------------- + ew = {} + x = torch.randn(1, L, HID, device=dev, dtype=torch.bfloat16) + br = torch.randn_like(x) + nw = torch.randn(HID, device=dev, dtype=torch.bfloat16) + tab = [torch.randn(3, HID, device=dev, dtype=torch.bfloat16) * 0.1 for _ in range(6)] + idx = torch.randint(0, 3, (L,), device=dev) + norm = torch.nn.RMSNorm(HID, eps=1e-6, device=dev, dtype=torch.bfloat16) + + def eager_mod(): + n1 = norm(x) * (1.0 + tab[1].index_select(0, idx)) + tab[0].index_select(0, idx) + h = x + tab[2].index_select(0, idx) * br + n2 = norm(h) * (1.0 + tab[4].index_select(0, idx)) + tab[3].index_select(0, idx) + out = h + tab[5].index_select(0, idx) * br + return n1, n2, out + + def fused_mod(): + n1 = fused_rmsnorm_modulate(x, nw, tab[1], tab[0], idx, 1e-6) + h, n2 = fused_residual_gate_rmsnorm_modulate(x, br, tab[2], nw, tab[4], tab[3], idx, 1e-6) + out = h + tab[5].index_select(0, idx) * br + return n1, n2, out + + with torch.no_grad(): + ew["modulate_eager_ms"] = timed(eager_mod) + ew["modulate_fused_ms"] = timed(fused_mod) + ew["modulate_compiled_ms"] = timed(torch.compile(eager_mod)) + packed = torch.randn(1, L, 2 * FFN, device=dev, dtype=torch.bfloat16) + + def eager_swiglu(): + a, b = packed.chunk(2, dim=-1) + return a * F.silu(b) + + ew["swiglu_eager_ms"] = timed(eager_swiglu) + ew["swiglu_fused_ms"] = timed(lambda: minimax_h3_swiglu(packed)) + del packed + qq = torch.randn(1, L, HEADS, HD, device=dev, dtype=torch.bfloat16) + cos = torch.randn(L, 96, device=dev, dtype=torch.bfloat16) + sin = torch.randn(L, 96, device=dev, dtype=torch.bfloat16) + qn = torch.nn.RMSNorm(HD, eps=1e-6, device=dev, dtype=torch.bfloat16) + + def eager_rope(): + outs = [] + for t in (qq, qq): + t = qn(t) + rot, pas = t[..., :96], t[..., 96:] + c, s = cos[None, :, None, :], sin[None, :, None, :] + a, b = rot.chunk(2, dim=-1) + outs.append(torch.cat((rot * c + torch.cat((-b, a), -1) * s, pas), -1).contiguous()) + return outs + + ew["qknorm_rope_eager_ms"] = timed(eager_rope) + ew["qknorm_rope_fused_ms"] = timed( + lambda: (fused_qknorm_rope(qq, qn.weight, cos, sin, 1e-6), fused_qknorm_rope(qq, qn.weight, cos, sin, 1e-6))) + r["elementwise"] = ew + report[name] = r + print("BENCH", name, json.dumps(r), flush=True) + return report + + +def check_tile64(): + """Exact semantics check of the quadrant path on small H3-style layouts.""" + from fastvideo.attention.backends.video_sparse_attn_h3 import _build_block_mask, _h3_tile_geometry, _pool_tiles + import attn_qat_infer.api as fa + from attn_qat_infer.api import sageattn_blackwell, sageattn_blackwell_sparse + + dev = torch.device("cuda") + gen = torch.Generator(device=dev).manual_seed(1) + out = {} + cases = [("odd_tiles", (77, 46), (9, 6, 10), 64), ("even_tiles", (64, 40), (8, 8, 8), 64), + ("tile256", (100, 70), (8, 8, 16), 256)] + for name, prefix, vshape, tt in cases: + H = 4 + shape = {64: (4, 4, 4), 256: (4, 8, 8)}[tt] + _, vbs, untile, n_prefix, n_video = _h3_tile_geometry(prefix, vshape, dev, shape) + n_tiles = vbs.numel() + Lpad = n_tiles * tt + L = sum(prefix) + math.prod(vshape) + q = torch.randn(1, L, H, HD, device=dev, generator=gen, dtype=torch.bfloat16) + k = torch.randn(1, L, H, HD, device=dev, generator=gen, dtype=torch.bfloat16) + v = torch.randn(1, L, H, HD, device=dev, generator=gen, dtype=torch.bfloat16) + + def tile(x): + buf = torch.zeros((1, Lpad, H, HD), device=dev, dtype=x.dtype) + buf[:, untile] = x + return buf.transpose(1, 2).contiguous() # BHSD + + qt, kt, vt = tile(q), tile(k), tile(v) + scores = torch.matmul(_pool_tiles(qt.transpose(1, 2), vbs, tt), _pool_tiles(kt.transpose(1, 2), vbs, tt).transpose(-2, -1)) + mask = _build_block_mask(scores, n_prefix, 0.8, True, ((n_prefix, n_prefix + n_video), ), (0.8, )) + q2k_idx, q2k_num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, tt, vbs, validate=True) + Lk = q2k_idx.shape[2] * 128 + pad = lambda x: F.pad(x, (0, 0, 0, Lk - Lpad)).contiguous() + o = sageattn_blackwell_sparse(pad(qt), pad(kt), pad(vt), q2k_idx, q2k_num, kvv, quad)[:, :, :Lpad] + tok_tile = torch.arange(n_tiles, device=dev).repeat_interleave(tt) + tok_valid = torch.zeros(Lpad, dtype=torch.bool, device=dev) + tok_valid[untile] = True + tm = mask[:, :, tok_tile][:, :, :, tok_tile] & tok_valid[None, None, None, :] + ref = F.scaled_dot_product_attention(qt.float(), kt.float(), vt.float(), attn_mask=tm) + rows = untile + err = (o.float() - ref)[:, :, rows] + # FP4 noise floor: dense FP4 vs dense fp32 on the same valid tokens + qv, kv_, vv = (x[:, :, rows].contiguous() for x in (qt, kt, vt)) + od = sageattn_blackwell(qv, kv_, vv) + rd = F.scaled_dot_product_attention(qv.float(), kv_.float(), vv.float()) + out[name] = { + "tiles": n_tiles, "density": round(mask.float().mean().item(), 3), + "rel_l2": round((err.norm() / ref[:, :, rows].norm()).item(), 4), + "dense_fp4_rel_l2_floor": round(((od.float() - rd).norm() / rd.norm()).item(), 4), + "cosine": round(F.cosine_similarity(o[:, :, rows].float().flatten(), ref[:, :, rows].flatten(), dim=0).item(), 5), + "finite": bool(torch.isfinite(o[:, :, rows]).all()), + } + if tt == 64: + from fastvideo_kernel.block_sparse_attn import block_sparse_attn as bsa64 + ot, _ = bsa64(qt, kt, vt, mask, vbs) + out[name]["triton_bf16_rel_l2"] = round(((ot.float() - ref)[:, :, rows].norm() / ref[:, :, rows].norm()).item(), 4) + print("CHECK", name, json.dumps(out[name]), flush=True) + return out + + +def density_study(prefix_segments=(256, 810), video_shape=(72, 24, 42)): + """Block densities each kernel granularity would compute, and FP4 sparse timing at 0.8 / 0.9.""" + from fastvideo.attention.backends.video_sparse_attn_h3 import _build_block_mask, _h3_tile_geometry, _pool_tiles + import attn_qat_infer.api as fa + + dev = torch.device("cuda") + gen = torch.Generator(device=dev).manual_seed(0) + _, vbs, untile, n_prefix, n_video = _h3_tile_geometry(tuple(prefix_segments), tuple(video_shape), dev, (4, 4, 4)) + n_tiles = vbs.numel() + Lpad = n_tiles * 64 + q = smooth_tokens(sum(prefix_segments), tuple(video_shape), dev, gen) + k = (q + 0.5 * torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16)).to(torch.bfloat16) + v = torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16) + + def tile(x): + buf = torch.zeros((1, Lpad, HEADS, HD), device=dev, dtype=x.dtype) + buf[:, untile] = x + return buf + + qt, kt, vt = tile(q), tile(k), tile(v) + scores = torch.matmul(_pool_tiles(qt, vbs, 64), _pool_tiles(kt, vbs, 64).transpose(-2, -1)) + nt2 = n_tiles + n_tiles % 2 + rows = nt2 * 64 + qs, ks, vs = (F.pad(x, (0, 0, 0, 0, 0, rows - Lpad)).contiguous() for x in (qt, kt, vt)) + out = {} + for sparsity in (0.8, 0.9): + mask = _build_block_mask(scores, n_prefix, sparsity, True, ((n_prefix, n_prefix + n_video), ), (sparsity, )) + m = F.pad(mask, (0, nt2 - n_tiles, 0, nt2 - n_tiles), value=False) + B, H = m.shape[:2] + d = {"tile64x64": m.float().mean().item()} + q2 = m.view(B, H, nt2 // 2, 2, nt2) # query pairs + d["q128_k64"] = q2.any(3).float().mean().item() + k2 = m.view(B, H, nt2, nt2 // 2, 2) + d["q64_k128"] = k2.any(4).float().mean().item() + d["q128_k128"] = m.view(B, H, nt2 // 2, 2, nt2 // 2, 2).any(5).any(3).float().mean().item() + idx, num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, 64, vbs) + def bshd(): + qh, kh, vh = (x.transpose(1, 2) for x in (qs, ks, vs)) + ds = fa._zero_delta_s(1, HEADS, rows, dev) + return fa.blockscaled_fp4_attn_sparse(fa.scale_and_quant_fp4(qh), fa.scale_and_quant_fp4_permute(kh), + fa.scale_and_quant_fp4_transpose(vh), ds, rows, idx, num, kvv, quad, + False, True, True, None)[0] + + d["sparse_fp4_quad_ms"] = timed(bshd, iters=3) + out[str(sparsity)] = {key: round(val, 4) for key, val in d.items()} + print("DENSITY", sparsity, json.dumps(out[str(sparsity)]), flush=True) + out["dense_fp4_bshd_ms"] = timed(lambda: fa.sageattn_blackwell(*(x.transpose(1, 2) for x in (qs, ks, vs))), iters=2) + return out diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py new file mode 100644 index 0000000000..cd25102905 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -0,0 +1,118 @@ +"""Headline e2e benchmark for the FastH3 local release: 480p 5 s and 768p 10 s. + +usage: python bench_headline.py [--settings 480p5s,768p10s] + +Protocol (RELEASE_PLAN §5): prompts latency-ceramics-005 + latency-harbor-005, one untimed warmup per +setting, then each prompt timed twice; e2e = generate_video wall time (encode + DiT + decode + mp4 write). +Per-stage times come from FASTVIDEO_STAGE_LOGGING. Clips and results.json go to //, +and one W&B run per invocation (group headline-) when WANDB_PROJECT is set. +""" +import argparse +import json +import os +import statistics +import time + +SETTINGS = {"480p5s": (832, 480, 124), "768p10s": (1344, 768, 243)} # 17n+5 frames at 24 fps +PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("run_name") + ap.add_argument("model_dir") + ap.add_argument("num_gpus", type=int) + ap.add_argument("nvfp4_profile") + ap.add_argument("--settings", default="480p5s,768p10s") + ap.add_argument("--prompts", default=os.path.join(os.path.dirname(os.path.abspath(__file__)), "headline_prompts.json")) + ap.add_argument("--out", default=os.environ.get("HEADLINE_OUT", "headline")) + ap.add_argument("--timed", type=int, default=2) + ap.add_argument("--sparsity", type=float, default=None, help="default: the checkpoint contract's vsa_sparsity") + a = ap.parse_args() + if a.timed < 1: + ap.error("--timed must be at least 1") + + contract = json.load(open(os.path.join(a.model_dir, "fastvideo_inference.json"))) + steps = contract["dmd_denoising_steps"] + sparsity = a.sparsity if a.sparsity is not None else float(contract.get("vsa_sparsity", 0.8)) + texts = json.load(open(a.prompts)) + out_dir = os.path.join(a.out, a.run_name) + os.makedirs(out_dir, exist_ok=True) + + engine = {"num_gpus": a.num_gpus, "use_fsdp_inference": False, + "parallelism": {"tp_size": 1, "sp_size": a.num_gpus}, + "offload": {"dit": False, "dit_layerwise": False, "text_encoder": False, "vae": False, + "pin_cpu_memory": False, "lazy_module_load": False}, + "compile": {"enabled": False, "vae_enabled": os.environ.get("HEADLINE_VAE_COMPILE", "1") == "1"}} + if os.environ.get("HEADLINE_BACKEND"): + engine["execution_backend"] = os.environ["HEADLINE_BACKEND"] # "ray" for multi-node + if a.nvfp4_profile != "none": + engine["quantization"] = {"transformer_quant": "NVFP4", "layer_profile": a.nvfp4_profile} + experimental = {"attention_backend": "VIDEO_SPARSE_ATTN_H3", "VSA_sparsity": sparsity, "VSA_tile_size": 64, + "h3_sequential_load": False, "inference_torch_compile": False} + if os.environ.get("HEADLINE_VAE_PARALLEL") == "1" and a.num_gpus > 1: + # Decode VAE tiles on every GPU instead of rank 0 only. + experimental.update(vae_parallel_decode=True, vae_parallel_decode_strategy="gather") + # Placement overrides for memory-limited GPUs, e.g. on a 32 GB RTX 5090: + # HEADLINE_ENGINE_JSON='{"offload": {"text_encoder": true, "pin_cpu_memory": true}}' + # HEADLINE_EXPERIMENTAL_JSON='{"h3_sequential_load": true}' + for key, value in json.loads(os.environ.get("HEADLINE_ENGINE_JSON", "{}")).items(): + if isinstance(value, dict) and isinstance(engine.get(key), dict): + engine[key].update(value) + else: + engine[key] = value + experimental.update(json.loads(os.environ.get("HEADLINE_EXPERIMENTAL_JSON", "{}"))) + config = {"model_path": a.model_dir, "engine": engine, "pipeline": {"experimental": experimental}} + # HEADLINE_* switches (VAE_PARALLEL, ENGINE_JSON, ...) change the run, so record them with the rest. + env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA", "HEADLINE_"))} + run = None + if os.environ.get("WANDB_PROJECT"): + import wandb + run = wandb.init(project=os.environ["WANDB_PROJECT"], entity=os.environ.get("WANDB_ENTITY"), + name=a.run_name, group=os.environ.get("HEADLINE_GROUP", "headline"), job_type="benchmark", + dir=out_dir, config={"model_dir": a.model_dir, "num_gpus": a.num_gpus, + "nvfp4_profile": a.nvfp4_profile, "dmd_steps": steps, + "vsa_sparsity": sparsity, "settings": a.settings, "engine": engine, + "experimental": experimental, + "env": env, "device": os.environ.get("HEADLINE_DEVICE", "")}) + + from fastvideo import VideoGenerator + t0 = time.perf_counter() + generator = VideoGenerator.from_config(config) + results = {"run_name": a.run_name, "model_dir": a.model_dir, "num_gpus": a.num_gpus, + "nvfp4_profile": a.nvfp4_profile, "load_s": round(time.perf_counter() - t0, 1), "env": env, + "engine": engine, "experimental": experimental, "settings": {}} + try: + for name in a.settings.split(","): + width, height, frames = SETTINGS[name] + runs = [] + plan = [(PROMPT_IDS[0], True)] + [(pid, False) for _ in range(a.timed) for pid in PROMPT_IDS] + for i, (pid, warmup) in enumerate(plan): + path = os.path.join(out_dir, f"{name}_{pid}_{'warmup' if warmup else i}.mp4") + t = time.perf_counter() + generator.generate_video(prompt=texts[pid], height=height, width=width, num_frames=frames, fps=24, + guidance_scale=1.0, num_inference_steps=len(steps) + 1, seed=1234, + output_path=path, save_video=True) + wall = round(time.perf_counter() - t, 2) + runs.append({"prompt": pid, "warmup": warmup, "e2e_s": wall, "path": path}) + print("RUN", name, pid, "warmup" if warmup else "timed", wall, flush=True) + if run is not None and not warmup: + import wandb + run.log({f"{name}/e2e_s": wall, f"{name}/{pid}": wandb.Video(path, fps=24, format="mp4")}) + timed = [r["e2e_s"] for r in runs if not r["warmup"]] + results["settings"][name] = {"width": width, "height": height, "frames": frames, + "e2e_median_s": statistics.median(timed), "e2e_min_s": min(timed), "runs": runs} + print("SETTING", name, json.dumps(results["settings"][name]), flush=True) + if run is not None: + run.summary[f"{name}_e2e_median_s"] = statistics.median(timed) + json.dump(results, open(os.path.join(out_dir, "results.json"), "w"), indent=1) + finally: + generator.shutdown() + if run is not None: + run.summary["load_s"] = results["load_s"] + run.finish() + print("HEADLINE_DONE", a.run_name, flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_pro6000/download_personal.py b/scripts/benchmarks/minimax_h3_pro6000/download_personal.py new file mode 100644 index 0000000000..eba4edbc55 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/download_personal.py @@ -0,0 +1,40 @@ +"""Stage every FastH3 input into the personal workspace's volume (CPU only, public repos, no token).""" +import modal + +volume = modal.Volume.from_name("h3-pro6000-weights", create_if_missing=True) +image = modal.Image.debian_slim(python_version="3.12").pip_install("huggingface_hub[hf_transfer]>=0.34") +app = modal.App("h3-pro6000-download-personal", image=image) + + +@app.function(volumes={"/vol": volume}, timeout=10800, cpu=16, memory=32768) +def download() -> dict: + import os + from huggingface_hub import hf_hub_download, snapshot_download + out = {} + out["v2"] = snapshot_download("FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4", local_dir="/vol/v2-nvfp4", max_workers=16) + volume.commit() + out["v4"] = snapshot_download("FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree-NVFP4", local_dir="/vol/v4-nvfp4", + allow_patterns=["transformer/*", "scheduler/*", "audio_scheduler/*", "*.json", + "tokenizer/*", "processor/*"], max_workers=16) + volume.commit() + out["official_vae"] = snapshot_download("MiniMaxAI/MiniMax-H3", local_dir="/vol/official", + allow_patterns=["vae/*", "model_index.json", "modular_model_index.json"], + max_workers=8) + out["comfy_int8_vae"] = hf_hub_download("Comfy-Org/MiniMax-H3", "vae/minimax_h3_video_vae_int8_convrot.safetensors", + local_dir="/vol/comfy") + out["light_vae"] = hf_hub_download("corechan/MiniMax-H3-LightVAE", "lynnreal_light_vae_decoder_fp16.safetensors", + local_dir="/vol/light-vae") + out["kijai"] = hf_hub_download("Kijai/MiniMax-H3-experimental", "minimax_h3_lynnreal_light_vae_int8_convrot.safetensors", + local_dir="/vol/kijai") + volume.commit() + sizes = {} + for root in ("/vol/v2-nvfp4", "/vol/v4-nvfp4", "/vol/official", "/vol/comfy", "/vol/light-vae", "/vol/kijai"): + sizes[root] = round(sum(os.path.getsize(os.path.join(d, f)) for d, _, fs in os.walk(root) for f in fs) / 1e9, 2) + out["sizes_gb"] = sizes + return out + + +@app.local_entrypoint() +def main(): + import json + print("DOWNLOADED", json.dumps(download.remote(), indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_prompts.json b/scripts/benchmarks/minimax_h3_pro6000/headline_prompts.json new file mode 100644 index 0000000000..553a5a7d67 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_prompts.json @@ -0,0 +1,4 @@ +{ + "latency-ceramics-005": "In a quiet pottery studio an adult potter steadies a small spinning bowl while an adult apprentice watches. The apprentice asks, \"Is the rim ready?\" The potter says, \"One more gentle pass,\" and smooths the lip with a damp sponge. Begin with a close view of the hands, then make one restrained cut to a shoulder-level view showing both faces. The wheel hum, damp clay, a light splash and breathing form the soundscape. The movement is careful and unhurried, with no background music. The entire event is one finishing pass on the same bowl, not a demonstration of the whole pottery process.\nThe apprentice's apron pocket is indigo, with a plain surface. This detail belongs to the existing object, stays at its established location and remains subordinate to the main action. Preserve its material and appearance through the camera movement.\nThe wheel sits near a tall north-facing window on the left wall of the studio.\nIts circular splash pan is at the seated potter's waist.\nThe apprentice stands beyond the right side of the wheel, with enough room between their clothes and the wet rim to avoid accidental contact.\nBehind the potter, a shallow shelf holds a few dry cups.\nA sink is farther back on the right, beneath a small rack of towels.\nThese positions remain fixed through the cut.\nThe wider view reveals the same work area that surrounds the hands in the opening close view.\nThe bowl is modest in size, comfortably held between two hands.\nIts side rises from a low foot in a continuous shallow curve and ends in a slightly thickened lip.\nThe clay is warm gray with a brown undertone, darker where it is wet.\nA narrow spiral line below the rim records the potter's earlier touch.\nThat line rotates with the bowl, while a stationary highlight from the window travels over the passing surface.\nKeep these two motions distinguishable: the material marks belong to the clay, and the reflected light belongs to the relationship between the surface and the window.\nOpen with the camera low enough to see the bowl's interior without looking directly down onto the wheel.\nThe potter's left fingertips support the inside edge.\nThe right hand holds a small natural sponge outside the lip.\nNeither hand blocks the whole form.\nThe near edge of the splash pan occupies the bottom of the composition as a soft curved boundary.\nBeyond the hands, the apprentice's apron is initially out of focus.\nThis arrangement gives the close shot depth and prepares the later view of the two people without requiring another establishing shot.\nThe left hand is already stable when the clip begins.\nIts fingers form a loose supportive curve rather than squeezing the wall.\nThe sponge approaches the outside edge with only a small adjustment of the wrist.\nAs it contacts the clay, it compresses slightly and darkens where moisture gathers.\nThe bowl continues to rotate at a steady moderate speed.\nShow the finishing pass as a change in the surface's smoothness and the evenness of the lip, not as a large change in the bowl's overall shape.\nThe work is nearly finished before this moment begins.\nThe apprentice asks the question while looking at the rim, then briefly lifts their eyes toward the potter.\nTheir hands rest loosely together in front of the apron, safely away from the rotating work.\nThe question is curious and quiet, with the natural upward inflection of someone checking a detail.\nThe potter answers without stopping the wheel or turning their whole body.\nA small glance toward the apprentice is sufficient before attention returns to the clay.\nKeep the spoken words exactly as given, with no narrator explaining the technique and no extra exchange after the answer.\nCut once after the question has made the apprentice's presence clear.\nThe shoulder-level view places the potter to the left and the apprentice to the right, preserving the established relation to the wheel.\nThe bowl remains visible between them in the lower part of the frame.\nThe potter's right hand still holds the same sponge at the same point on the rim.\nContinue the wheel sound across the cut without a restart.\nThe change of view should feel like a closer understanding of the same instant, not a jump forward to a later stage of the work.\nThe potter wears a practical cotton work shirt with the sleeves rolled above the wrists.\nThe folds gather at the elbows and remain dry there.\nSmall clay marks on the forearms and apron are concentrated near the work area.\nThey do not spread or migrate during the pass.\nThe apprentice's apron is cleaner but shows a few dry pale smudges near one pocket.\nBoth garments have weight and ordinary creases.\nAvoid pristine costumes or exaggerated distressing; this is a used studio where people work carefully and clean their tools regularly.\nGive the potter a focused, patient expression.\nTheir mouth moves only for the brief reply, then settles while they feel the rim through the sponge.\nAllow the last small movement to settle within the established composition, with the environmental sound continuing around it. Preserve the quiet final composition.", + "latency-harbor-005": "## Harbor: the tide chart\nThe fabric cover around the folded tide chart is indigo, with a plain surface. This detail belongs to the existing object, stays at its established location and remains subordinate to the main action. Preserve its material and appearance through the camera movement.\nThe event takes place beside a small passenger ferry tied to a working harbor pier just before sunrise. An adult mechanic stands on the ferry's open side deck, and its captain stands beside the cabin entrance. The mechanic offers a folded tide chart and says, \"The channel is clear.\" The captain accepts it, answers, \"Then we can go,\" and looks out toward the harbor entrance. A slow lateral camera move reveals the channel beyond their shoulders. The boat remains moored throughout this brief exchange. Close voices, water against the hull, a loose halyard and a distant gull make the soundscape; there is no music.\nThe ferry is a practical coastal launch with a dark blue hull and a narrow cream band beneath its windows.\nIt carries a small enclosed cabin forward and an open passenger area behind it.\nThe camera is on the open deck, looking diagonally toward the cabin and the gap between the two people.\nThis angle places the pier along the left edge of the view and open water farther to the right.\nThe horizon is low enough that the upper part of the cabin has a clear silhouette against the pale sky.\nNothing in the composition suggests that the ferry is already underway.\nThe mechanic has finished a routine inspection rather than an emergency repair.\nTheir expression is alert but comfortable, with the slight tiredness of an early start.\nThey wear a plain work jacket over a warm shirt and carry no conspicuous badge or brand.\nA few old creases in the jacket show where the elbows bend.\nThe sleeve nearest the chart has a darker damp patch near its cuff from resting against the rail.\nKeep that patch in the same place as the arm moves.\nThe mechanic's free hand rests lightly on the top of a closed tool bag at hip level.\nThe captain is a different adult, dressed for a cool morning outside.\nA heavy knit sweater is visible beneath an open weatherproof coat.\nTheir hair is tidy but not freshly styled, and the light catches a few loose strands when they turn toward the water.\nTheir stance is balanced on the gently moving deck, with one foot slightly ahead of the other.\nThey are listening to the mechanic before the first line begins.\nThe captain does not interrupt or make a broad theatrical gesture.\nTheir reply is a small decision shared between people accustomed to working together.\nThe tide chart is a real paper object with several old folds.\nIt is partly folded into a rectangle that can be held in one hand, but one narrow flap remains loose.\nFaint printed lines and numbers are visible as a texture on its pale surface without becoming a readable title or a map inset.\nA soft graphite mark near one fold suggests that it has been used for planning.\nThe mechanic holds its lower edge between the thumb and fingers, keeping the paper clear of the damp rail.\nIts upper corner lifts slightly in the breeze before the captain takes it.\nBegin with both people already in the frame.\nThe mechanic's hand and the chart occupy the space between their bodies, below their faces.\nThis arrangement lets the first line and the handover belong to the same view.\nAs the mechanic speaks, the chart moves a short distance toward the captain.\nThe motion is neither a flourish nor an abrupt thrust.\nThe captain's receiving hand rises from beside the coat, touches the opposite edge and supports it before the mechanic releases their grip.\nThe paper bends a little between the two hands during that shared moment of support.\nThe mechanic's line is spoken in an ordinary low voice suitable for the quiet morning. The consonants remain clear, but the delivery does not sound like a public announcement. Their mouth and jaw make the small movements of the exact words, and their eyes remain on the captain. There is a slight release of breath after \"clear.\" The captain acknowledges the information with a very small nod before replying. The pause is long enough to register listening and short enough that the exchange feels familiar. Do not add another question, greeting or explanation of the voyage.\nWhen the captain says, \"Then we can go,\" the first part of the line is addressed to the mechanic. On the last words, their gaze begins to move toward the channel. The head follows the eyes through a modest turn, revealing more of the cheek nearest the exterior light. The chart settles against the front of the coat, still visibly held. The mechanic follows the captain's look with a quieter change of attention. Both remain in place. The ending is anticipation of departure, not departure itself: no engine surge, released rope or sudden movement of the ferry is needed.\nThe camera makes a restrained lateral movement toward the open-water side of the deck.\nKeep the final gesture restrained and preserve the surrounding atmosphere. Preserve the quiet final composition as the scene reaches its stated resolution." +} \ No newline at end of file diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py new file mode 100644 index 0000000000..463a8371c8 --- /dev/null +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -0,0 +1,261 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Convert a ModelOpt-unified NVFP4 MiniMax-H3 transformer into FastVideo's packed export. + +ModelOpt's unified Hugging Face layout (e.g. ``FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4``) +stores each quantized linear as:: + + .weight uint8 [out, in // 2] two E2M1 values per byte + .weight_scale float8_e4m3fn [out, in // 16] linear layout + .weight_scale_2 float32 [] amax(|W|) / (6 * 448) + .input_scale float32 [] static activation scale + +FastVideo's packed NVFP4 DiT export (``nvfp4_weights.safetensors``, read by +``load_minimax_h3_nvfp4_dit_export``) stores ``::_nvfp4_weight`` (same +bytes), ``::_nvfp4_weight_scale`` (the same E4M3 bytes in FlashInfer's 128x4 +swizzled layout), ``::_nvfp4_alpha`` (= ``weight_scale_2``) and +``::_weight_global_sf`` (= 1 / ``weight_scale_2``), and +``::_nvfp4_input_global_sf`` (= 1 / ``input_scale``). The calibrated weight +bytes and activation scale are preserved. Dropping the activation scale would +replace its calibrated range with a unit global scale, clipping inputs above +2688. + +``--quantize-attention`` additionally quantizes the dense BF16 attention +projections (``attn.to_{q,k,v,out}``) of every main block exactly as +``convert_model_to_nvfp4`` would at runtime, producing the full +``layer_profile="h3_dit"`` set. ``--quantize-gate`` (with it) also quantizes +each block's VSA compression gate, for ``layer_profile="h3_dit_vsa"``. +Without either the export holds the FFN linears only and must be loaded with +``layer_profile="h3_dit_ffn"``. + +``--quantize-ffn`` takes the FFN linears from a dense BF16 source instead of a +ModelOpt export, with the same round-to-nearest weight math as the runtime +(ModelOpt's max calibration produces the same weight codes). ``--act-amax`` +adds a calibrated static activation scale per linear +(``_nvfp4_input_global_sf`` = 448 * 6 / amax) from a JSON of input amax keyed +``b.`` (e.g. ``b3.ff.fc_in``); without it activations use the +unit global scale, which saturates inputs above 2688 (H3's ``ff.fc_out``). + +Every exported linear is probed on random BF16 rows through the same +``mm_fp4`` path the loader runs; the relative error against a BF16 matmul with +the dequantized weight must stay under ``--max-probe-error`` (genuine W4A4 +noise on random inputs is about 0.1; a wrong nibble order or scale layout +reads near 1.0). Needs a Blackwell GPU with FlashInfer. + +Usage:: + + python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \\ + --src /path/to/FastH3-V2-NVFP4/transformer --dst /path/to/out/transformer [--quantize-attention] +""" +from __future__ import annotations + +import argparse +import json +import re +import shutil +from pathlib import Path + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +EXPORT_FILENAME = "nvfp4_weights.safetensors" +# ModelOpt keys not copied as dense tensors. ``input_scale`` (ModelOpt's static activation scale) +# is carried into the packed export as ``_nvfp4_input_global_sf`` = 1 / input_scale, unless +# ``--act-amax`` supplies a replacement calibration. +DROPPED_MODELOPT_SUFFIXES = ("input_scale",) +_BLOCK_ATTN = re.compile(r"^transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out\.0)$") +_BLOCK_FFN = re.compile(r"^transformer_blocks\.\d+\.ff\.net\.(?:0\.proj|2)$") +_BLOCK_GATE = re.compile(r"^transformer_blocks\.\d+\.attn\.to_gate_compress$") +_RENAMES = ((re.compile(r"\.ff\.net\.0\.proj$"), ".ff.fc_in"), (re.compile(r"\.ff\.net\.2$"), ".ff.fc_out"), + (re.compile(r"\.attn\.to_out\.0$"), ".attn.to_out")) +_E2M1 = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0]) + + +def fastvideo_module_name(diffusers_prefix: str) -> str: + name = diffusers_prefix + for pattern, replacement in _RENAMES: + name = pattern.sub(replacement, name) + return name + + +def dequantize_modelopt(packed: torch.Tensor, scale: torch.Tensor, scale_2: torch.Tensor) -> torch.Tensor: + """E2M1 low nibble = even column, high nibble = odd column (ModelOpt / CUTLASS order).""" + lut = _E2M1.to(packed.device) + low = lut[(packed & 0x0F).long()] + high = lut[(packed >> 4).long()] + values = torch.stack((low, high), dim=-1).reshape(packed.shape[0], packed.shape[1] * 2) + block = scale.to(torch.float32).repeat_interleave(16, dim=1) + return values * block * scale_2.to(torch.float32) + + +def _flashinfer(): + from flashinfer import SfLayout, mm_fp4, nvfp4_quantize + try: + from flashinfer import block_scale_interleave + except ImportError: + from flashinfer.fp4_quantization import block_scale_interleave + return SfLayout, mm_fp4, nvfp4_quantize, block_scale_interleave + + +def probe(buffers: dict[str, torch.Tensor], reference: torch.Tensor, rows: int = 256) -> float: + """Relative error of the loader's mm_fp4 path against BF16 x @ W_ref^T.""" + SfLayout, mm_fp4, nvfp4_quantize, _ = _flashinfer() + device = buffers["_nvfp4_weight"].device + x = torch.randn(rows, reference.shape[1], device=device, dtype=torch.bfloat16) + unit = torch.tensor(1.0, device=device, dtype=torch.float32) + x_fp4, x_scale = nvfp4_quantize(x, unit, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + out = mm_fp4(x_fp4, buffers["_nvfp4_weight"].T, x_scale, buffers["_nvfp4_weight_scale"].T, + buffers["_nvfp4_alpha"] / unit, torch.bfloat16, None, backend="auto") + ref = x.float() @ reference.float().T + return ((out.float() - ref).norm() / ref.norm()).item() + + +def convert_modelopt_linear(weight, scale, scale_2, device, input_scale=None + ) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: + """Carry the calibrated bytes over; return (buffers, dequantized weight, scale-byte agreement). + + The agreement compares the swizzled ModelOpt scales with the scales + ``nvfp4_quantize`` derives from the dequantized weight. It is a layout check + (identical swizzle and shape) and is near 1.0 when ModelOpt used max + calibration; blocks whose largest code is below 6 legitimately differ. + """ + SfLayout, _, nvfp4_quantize, block_scale_interleave = _flashinfer() + weight = weight.to(device) + scale = scale.to(device) + scale_2 = scale_2.to(device=device, dtype=torch.float32).reshape(()) + reference = dequantize_modelopt(weight, scale, scale_2) + _, layout_ref = nvfp4_quantize(reference.to(torch.bfloat16), 1.0 / scale_2, sfLayout=SfLayout.layout_128x4, + do_shuffle=False) + swizzled = block_scale_interleave(scale.view(torch.uint8).contiguous()).reshape(layout_ref.shape) + agreement = (swizzled.view(torch.uint8) == layout_ref.view(torch.uint8)).float().mean().item() + buffers = { + "_nvfp4_weight": weight.contiguous(), + "_nvfp4_weight_scale": swizzled.view(layout_ref.dtype).contiguous(), + "_nvfp4_alpha": scale_2.clone(), + "_weight_global_sf": (1.0 / scale_2).to(torch.bfloat16), + } + if input_scale is not None: + value = input_scale.to(dtype=torch.float32) + if value.numel() != 1 or not bool(torch.isfinite(value).all()) or not bool((value > 0).all()): + raise ValueError("ModelOpt input_scale must be a finite positive scalar") + buffers["_nvfp4_input_global_sf"] = value.reshape(()).reciprocal().to(device=device) + return buffers, reference, agreement + + +def quantize_dense_linear(weight: torch.Tensor, device) -> tuple[dict[str, torch.Tensor], torch.Tensor]: + """Same math as ``nvfp4_config.convert_model_to_nvfp4``.""" + SfLayout, _, nvfp4_quantize, _ = _flashinfer() + weight = weight.to(device=device, dtype=torch.bfloat16) + global_sf = (448 * 6) / weight.float().abs().nan_to_num().max() + fp4_w, fp4_s = nvfp4_quantize(weight, global_sf, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + global_sf = torch.as_tensor(global_sf, device=device, dtype=torch.float32) + buffers = {"_nvfp4_weight": fp4_w, "_nvfp4_weight_scale": fp4_s, "_nvfp4_alpha": (1.0 / global_sf).float(), + "_weight_global_sf": global_sf.to(torch.bfloat16)} + return buffers, weight + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--src", required=True, type=Path) + parser.add_argument("--dst", required=True, type=Path) + parser.add_argument("--quantize-attention", action="store_true") + parser.add_argument("--quantize-gate", action="store_true", help="also quantize attn.to_gate_compress") + parser.add_argument("--quantize-ffn", action="store_true", help="quantize the FFN linears from a bf16 source") + parser.add_argument("--act-amax", type=Path, help="JSON of calibrated input amax per linear") + parser.add_argument("--max-probe-error", type=float, default=0.3) + parser.add_argument("--dense-shard-gb", type=float, default=5.0) + args = parser.parse_args() + device = torch.device("cuda") + args.dst.mkdir(parents=True, exist_ok=True) + + index = json.loads((args.src / "diffusion_pytorch_model.safetensors.index.json").read_text()) + weight_map: dict[str, str] = index["weight_map"] + modelopt = sorted(k[:-len(".weight_scale_2")] for k in weight_map if k.endswith(".weight_scale_2")) + modelopt_keys = {f"{p}.{s}" for p in modelopt for s in ("weight", "weight_scale", "weight_scale_2", *DROPPED_MODELOPT_SUFFIXES)} + if args.quantize_gate and not args.quantize_attention: + parser.error("--quantize-gate requires --quantize-attention") + selected = ([_BLOCK_ATTN.pattern] if args.quantize_attention else []) + ( + [_BLOCK_GATE.pattern] if args.quantize_gate else []) + ([_BLOCK_FFN.pattern] if args.quantize_ffn else []) + dense_pattern = re.compile("|".join(selected)) if selected else None + # Projections ModelOpt already quantized stay on the ModelOpt path; re-quantizing their packed + # uint8 weight as if it were dense BF16 would corrupt them. + modelopt_set = set(modelopt) + attention = sorted(k[:-len(".weight")] for k in weight_map + if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")]) + and k[:-len(".weight")] not in modelopt_set) if dense_pattern else [] + if args.quantize_ffn and any(_BLOCK_FFN.match(p) for p in modelopt): + parser.error("--quantize-ffn needs a bf16 source; this one already holds ModelOpt FFN linears") + amax_table = None + if args.act_amax: + raw = json.loads(args.act_amax.read_text()) + amax_table = {k: float(v["all"] if isinstance(v, dict) else v) for k, v in raw.items()} + attention_keys = {f"{p}.weight" for p in attention} + + readers = {shard: safe_open(str(args.src / shard), framework="pt", device="cpu") for shard in set(weight_map.values())} + tensor = lambda key: readers[weight_map[key]].get_tensor(key) + + export: dict[str, torch.Tensor] = {} + worst = 0.0 + agreements: list[float] = [] + attention_set = set(attention) + for prefix in modelopt + attention: + if prefix in attention_set: + buffers, reference = quantize_dense_linear(tensor(f"{prefix}.weight"), device) + else: + buffers, reference, agreement = convert_modelopt_linear(tensor(f"{prefix}.weight"), + tensor(f"{prefix}.weight_scale"), + tensor(f"{prefix}.weight_scale_2"), device, + input_scale=tensor(f"{prefix}.input_scale") + if f"{prefix}.input_scale" in weight_map else None) + agreements.append(agreement) + error = probe(buffers, reference) + worst = max(worst, error) + if error > args.max_probe_error: + raise SystemExit(f"probe error {error:.3f} on {prefix} exceeds {args.max_probe_error}; nothing written") + module = fastvideo_module_name(prefix) + if amax_table is not None: + block = re.match(r"transformer_blocks\.(\d+)\.(.+)$", module) + key = f"b{block.group(1)}.{block.group(2)}" + if key not in amax_table: + raise SystemExit(f"--act-amax has no entry {key!r} for {module}; nothing written") + buffers["_nvfp4_input_global_sf"] = torch.tensor((448.0 * 6.0) / max(amax_table[key], 1e-12), + dtype=torch.float32) + for name, value in buffers.items(): + export[f"{module}::{name}"] = value.cpu() + save_file(export, str(args.dst / EXPORT_FILENAME)) + + dense_keys = [k for k in weight_map if k not in modelopt_keys and k not in attention_keys] + shard, shard_bytes, shards = {}, 0, [] + for key in dense_keys: + value = tensor(key) + shard[key] = value + shard_bytes += value.numel() * value.element_size() + if shard_bytes >= args.dense_shard_gb * 1e9: + shards.append(shard) + shard, shard_bytes = {}, 0 + if shard: + shards.append(shard) + new_map = {} + for i, part in enumerate(shards, 1): + name = f"diffusion_pytorch_model-{i:05d}-of-{len(shards):05d}.safetensors" + save_file(part, str(args.dst / name)) + new_map.update({k: name for k in part}) + (args.dst / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps({"metadata": {}, "weight_map": new_map}, indent=2) + "\n") + config = json.loads((args.src / "config.json").read_text()) + config.pop("quantization_config", None) + (args.dst / "config.json").write_text(json.dumps(config, indent=2) + "\n") + for extra in args.src.iterdir(): + if extra.suffix not in (".safetensors", ".json"): + shutil.copy2(extra, args.dst / extra.name) + print(json.dumps({"exported_linears": len(modelopt) + len(attention), "modelopt_linears": len(modelopt), + "quantized_dense_linears": len(attention), "worst_probe_error": round(worst, 4), + "static_activation_scales": sum(k.endswith("::_nvfp4_input_global_sf") for k in export), + "gate_linears": sum(1 for p in attention if _BLOCK_GATE.match(p)), + "min_scale_byte_agreement": round(min(agreements), 4) if agreements else None, + "mean_scale_byte_agreement": round(sum(agreements) / len(agreements), 4) if agreements else None, + "dense_tensors": len(dense_keys), "dense_shards": len(shards)})) + + +if __name__ == "__main__": + main() diff --git a/tests/local_tests/benchmarks/test_h3_benchmark_memory.py b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py new file mode 100644 index 0000000000..72e41829c2 --- /dev/null +++ b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py @@ -0,0 +1,137 @@ +"""Benchmark capacity sampling supports both pod cgroup versions, without CUDA.""" +import importlib.util +import json +from pathlib import Path +import sys +from types import SimpleNamespace + +import pytest + + +SCRIPT = Path(__file__).resolve().parents[3] / "scripts/benchmarks/minimax_h3_4090/bench_pod.py" +SPEC = importlib.util.spec_from_file_location("h3_benchmark_memory", SCRIPT) +BENCH = importlib.util.module_from_spec(SPEC) +SPEC.loader.exec_module(BENCH) + + +class OneSample: + def __init__(self): + self.finished = False + + def is_set(self): + return self.finished + + def wait(self, _seconds): + self.finished = True + + +def sample(root, monkeypatch): + monkeypatch.setitem(sys.modules, "pynvml", SimpleNamespace( + nvmlInit=lambda: None, nvmlShutdown=lambda: None, + nvmlDeviceGetHandleByIndex=lambda _index: None, + nvmlDeviceGetMemoryInfo=lambda _handle: SimpleNamespace(used=8192))) + peak = BENCH.HostMemoryPeak(root) + peak.stop = OneSample() + peak._sample() + return peak + + +@pytest.mark.parametrize("version", [1, 2]) +def test_cgroup_memory_excludes_cached_files_from_anonymous_peak(tmp_path, monkeypatch, version): + if version == 2: + (tmp_path / "memory.current").write_text("4096") + (tmp_path / "memory.stat").write_text("anon 1024\nfile 3072\n") + else: + root = tmp_path / "memory" + root.mkdir() + (root / "memory.usage_in_bytes").write_text("4096") + (root / "memory.stat").write_text("rss 512\ntotal_rss 1024\ntotal_cache 3072\n") + peak = sample(tmp_path, monkeypatch) + assert peak.host_error is None + assert peak.peak_bytes == 4096 + assert peak.peak_anon_bytes == 1024 + assert peak.peak_gpu_bytes == 8192 + + +def test_missing_host_counters_do_not_disable_gpu_capacity_measurement(tmp_path, monkeypatch): + peak = sample(tmp_path, monkeypatch) + assert peak.host_error is not None + assert peak.peak_gpu_bytes == 8192 + + +def test_malformed_host_counters_do_not_report_a_measured_zero(tmp_path, monkeypatch): + (tmp_path / "memory.current").write_text("invalid") + peak = sample(tmp_path, monkeypatch) + assert peak.host_error is not None + assert peak.peak_gpu_bytes == 8192 + assert peak.metrics()["peak_host_cgroup_gib"] is None + assert peak.metrics()["peak_host_anon_gib"] is None + + +@pytest.mark.parametrize("fail,once", [(True, False), (False, False), (False, True)]) +def test_generation_receipts_preserve_memory_and_once_never_adds_clips(tmp_path, monkeypatch, fail, once): + model = tmp_path / "model" + model.mkdir() + (model / "fastvideo_inference.json").write_text("{}") + prompts = tmp_path / "prompts.json" + prompts.write_text('{"ceramics": "test"}') + shutdown = [] + requests = [] + + class FailingGenerator: + def generate(self, _request): + requests.append(_request) + if fail: + raise RuntimeError("CUDA out of memory") + + def shutdown(self): + shutdown.append(True) + + class Memory: + def __enter__(self): + return self + + def __exit__(self, *_args): + pass + + def metrics(self): + return {"peak_gpu_used_gib": 11.9, "peak_host_cgroup_gib": 28.3, + "peak_host_anon_gib": 20.1, "host_memory_error": None} + + fake_torch = SimpleNamespace(__version__="test", version=SimpleNamespace(cuda="test"), + cuda=SimpleNamespace(get_device_name=lambda _index: "test GPU")) + fake_video = SimpleNamespace(VideoGenerator=SimpleNamespace(from_config=lambda _config: FailingGenerator())) + monkeypatch.setitem(sys.modules, "torch", fake_torch) + monkeypatch.setitem(sys.modules, "fastvideo", fake_video) + monkeypatch.setattr(BENCH, "HostMemoryPeak", Memory) + monkeypatch.setattr(BENCH.subprocess, "check_output", lambda *_args, **_kwargs: "test GPU") + argv = [str(SCRIPT), "failure", str(model), "fp8", "--prompt-file", str(prompts), + "--output-root", str(tmp_path / "outputs")] + if once: + argv.append("--once") + monkeypatch.setattr(sys, "argv", argv) + if fail: + with pytest.raises(RuntimeError, match="CUDA out of memory"): + BENCH.main() + else: + BENCH.main() + raw = json.loads((tmp_path / "outputs/failure/results.json").read_text()) + if fail: + assert raw["runs"] == [] + assert len(raw["failed_runs"]) == 1 + failed = raw["failed_runs"][0] + assert failed["peak_gpu_used_gib"] == 11.9 + assert failed["peak_host_cgroup_gib"] == 28.3 + assert failed["warmup"] + assert "CUDA out of memory" in failed["error"] + else: + assert len(requests) == len(raw["runs"]) == (1 if once else 3) + assert raw["runs"][0]["warmup"] + assert all(run["peak_gpu_used_gib"] == 11.9 for run in raw["runs"]) + if once: + assert "median_e2e_s" not in raw + assert "mean_e2e_s" not in raw + else: + assert not any(run["warmup"] for run in raw["runs"][1:]) + assert "mean_e2e_s" in raw + assert shutdown == [True]