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 8875361dcc..19839b8b7d 100644 --- a/docs/assets/cookbook-recipes.json +++ b/docs/assets/cookbook-recipes.json @@ -1,5 +1,5 @@ { - "version": 11, + "version": 13, "recipes": [ { "id": "fastwan21-t2v", @@ -578,6 +578,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/cookbook/cosmos.md b/docs/cookbook/cosmos.md index 2ee864e5b7..89e39e8457 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 b8bbb59052..82508bd81d 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 b4a2c39673..51f4a7f28e 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 8408dad626..75caf3144d 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 e07993f24d..bbd96547b5 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 f898060525..3b43788d43 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 de1bf1c982..53bf0fa14f 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 8bcb544f55..15b62ae34a 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 8f3b4f33b2..e664054eac 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 8bdd93515a..a6e5654e77 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 a2ca588ef9..b2ccbc533b 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 a3c8c62ab2..653b636779 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 92b259fec4..860abee1cf 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 4554e09660..a7a55d0009 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 8f4301f647..7f39034dc8 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/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..03ea3a6565 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,166 @@ 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 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): + """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]``. + """ + QL = q.size(2) + 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) + 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) -> 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}") + 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 +371,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/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/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 05a9aacf54..1356219b2f 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -1,4 +1,5 @@ from contextlib import contextmanager +from itertools import chain from typing import Any import torch from torch import nn @@ -14,6 +15,30 @@ 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. + """ + import os + 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 os.environ.get("FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS") == "1": + 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__( @@ -35,7 +60,7 @@ 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(): + 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()) param.data = _tensor_placeholder(param.data, self.device) @@ -44,7 +69,7 @@ def on_init(self, module: nn.Module): 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 +81,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 +91,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,7 +108,7 @@ 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()) + 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: @@ -142,9 +167,20 @@ 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. + import os + try: + resident = max(0, int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0"))) + except ValueError: + logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%r", + os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS")) + resident = 0 for name, submodule in model.named_children(): if isinstance(submodule, nn.ModuleList): for idx, module_entry in enumerate(submodule): + if idx < resident: + continue state = LayerwiseOffloadState(async_copy_stream=async_stream, device=device) state_list.append(state) hook_mgr = ModuleHookManager.get_from_or_default(module_entry) @@ -159,6 +195,10 @@ 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 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..ded48937fa 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 @@ -76,6 +78,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 +101,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 +340,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 os.environ.get("FASTVIDEO_NVFP4_MM_BACKEND", "auto") + + def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: """Coerce an activation to a dtype the FP4 linear accepts. @@ -306,6 +366,18 @@ 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): def __init__(self, layer_prefix: str = ""): @@ -333,6 +405,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 = os.environ.get("FASTVIDEO_NVFP4_ACT_AMAX") + 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 = os.environ.get("FASTVIDEO_NVFP4_DYNAMIC_ACT", "") + 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 +500,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 +535,7 @@ def apply( alpha, torch.bfloat16, None, - backend='auto', + backend=_mm_fp4_backend(), ) if bias is not None: @@ -431,14 +549,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 +599,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 os.environ.get("FASTVIDEO_H3_FP8_ATTENTION") == "1" + 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=os.environ.get("FASTVIDEO_H3_FP8_GRANULARITY", "tensor")) return None @@ -534,9 +673,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 +691,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/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index e25650466d..ccda038983 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -4,6 +4,8 @@ from __future__ import annotations import contextlib +import json +import os import math from collections.abc import Iterator, Mapping from typing import Any @@ -14,6 +16,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 +25,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 +36,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) from fastvideo.models.dits.minimax_h3_fusions import ( HAVE_TRITON, fused_qknorm_rope, @@ -74,6 +82,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 +140,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 = int(os.environ.get("FASTVIDEO_H3_FFN_CHUNK_TOKENS", "0")) 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 +238,10 @@ 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.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,6 +321,21 @@ def forward( rotary_emb: tuple[torch.Tensor, torch.Tensor] | None, original_seq_len: int, ) -> torch.Tensor: + 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 query, _ = self.to_q(hidden_states) key, _ = self.to_k(hidden_states) value, _ = self.to_v(hidden_states) @@ -421,7 +475,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 os.environ.get("FASTVIDEO_H3_ADALN_DUMP"): + # 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)) @@ -564,8 +662,10 @@ 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"): + 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 +991,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 = os.environ.get("FASTVIDEO_H3_ADALN_DUMP") + 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 +1079,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 +1125,17 @@ 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}"): + with nvtx_range(f"minimax_h3.transformer_block.{block_index}"), STAGES.span("block_total"): packed_hidden_states = block( packed_hidden_states, temb, @@ -997,6 +1157,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..5429f76a3f --- /dev/null +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -0,0 +1,364 @@ +# 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 +import os +from typing import Any + +import torch + +from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAMetadata, _build_block_mask, _pool_tiles) + +VSA_FP4_ENV = "FASTVIDEO_H3_VSA_FP4" +_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 = os.environ.get("FASTVIDEO_H3_SP_PROFILE", "0") == "1" + 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 os.environ.get(VSA_FP4_ENV, "0") == "1" + + +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]: + """Run projections of one input, quantizing it once when all are NVFP4 with the unit activation scale. + + Only then is one quantized copy exactly what each layer would have produced; layers with a calibrated + or dynamic activation scale quantize their own input. + """ + from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod + + methods = [linear.quant_method for linear in linears] + if not 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)): + 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_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 = os.environ.get(SIM_SP_FP8_ENV, "0") == "1" + 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"): + out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) + 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. +SIM_SP_FP8_ENV = "FASTVIDEO_H3_SIM_SP_FP8" + + +@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) + + 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..686346e8ba 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -54,6 +54,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 +185,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 +302,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 " @@ -385,6 +425,12 @@ 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. + weight = dequantize_serialized_nvfp4(layer.weight_packed, layer.weight_scale, + float(layer.weight_global_scale.item()), 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. diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 148a6d1bb0..7e668a1709 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -469,7 +469,13 @@ def load_model( 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 +882,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 +925,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 +1085,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 +1139,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 = os.environ.get("FASTVIDEO_H3_ADALN_TABLE") or None + adaln_host_cache = (fastvideo_args.inference_mode and not fastvideo_args.use_fsdp_inference + and (os.environ.get("FASTVIDEO_H3_ADALN_CACHE") == "1" 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 +1156,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 +1190,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) + # 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 = os.environ.get("FASTVIDEO_H3_SPLICE_TRANSFORMER") + 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 = int(os.environ.get("FASTVIDEO_H3_SPLICE_FROM_STEP", "4")) + 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..01ccf3a75a 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -6,6 +6,7 @@ from __future__ import annotations import contextlib +import os import re from collections.abc import Callable, Generator from itertools import chain @@ -36,6 +37,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 +124,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 +274,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 +351,19 @@ 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 + load_weights_to_cpu = 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 +371,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 +379,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 os.environ.get("FASTVIDEO_H3_ADALN_TABLE"): + # 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 +402,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 +685,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 +702,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 +741,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 +813,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..649541fd19 --- /dev/null +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -0,0 +1,317 @@ +# 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 + +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.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 forward(self, x: torch.Tensor) -> torch.Tensor: + original_shape = x.shape + x_2d = x.reshape(-1, original_shape[-1]).contiguous() + if self.convrot: + x_2d = rotate_activation(x_2d, self.group_size) + if x_2d.device.type == "cuda" and x_2d.shape[-1] % 8 == 0: + 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) + # torch._int_mm requires M > 16. VAE decode is far above that; + # pad only the leftover short rows. + 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)) + acc = torch._int_mm(x_q, self.weight.t().contiguous())[:rows] + x_scale = x_scale[:rows] + out = self._dequant_int8_gemm(acc, x_scale, self.weight_scale) + else: + 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 _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_video.py b/fastvideo/models/vaes/minimax_h3_video.py index e116defb7f..21be0864a9 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -7,6 +7,7 @@ """ import math +import os from collections.abc import Iterator from dataclasses import dataclass @@ -522,6 +523,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, int(os.environ.get("FASTVIDEO_H3_VAE_TILE_BATCH", "1"))) + + 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 +823,28 @@ 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) + decoded.extend(self.decoder(self._project_decoder_tile(batch)).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..c09e39b0e9 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -1,9 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 -"""FastVideo composed pipelines for MiniMax H3.""" from __future__ import annotations import gc +import os import json import math from dataclasses import dataclass @@ -11,6 +11,7 @@ from typing import Any import torch +from torch.distributed.tensor import DTensor from fastvideo.attention.selector import (_active_component_attention_backend_scope, coerce_attn_backend, get_env_variable_attn_backend) @@ -29,6 +30,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 +38,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 +76,75 @@ def _checkpoint_has_vsa_gates(transformer_dir: Path) -> bool: return False +def _exact_pinned_views(tensors: list[torch.Tensor]) -> tuple[list[torch.Tensor], torch.Tensor | None]: + """Page-locked host copies of ``tensors`` backed by one exact-size allocation. + + ``Tensor.pin_memory()`` goes through torch's caching host allocator, which rounds every block up to a power + of two (1.76x for H3's packed FFN weights), so pinning a 20 GB DiT plus a 15 GB encoder overruns a 60 GB + container. Registering one plain allocation with ``cudaHostRegister`` pins exactly what is needed; the views + stay pinned and keep the arena alive. + """ + sizes = [-(-t.numel() * t.element_size() // 256) * 256 for t in tensors] + arena = torch.empty(max(sum(sizes), 1), dtype=torch.uint8) + cudart = torch.cuda.cudart() + if cudart.cudaHostRegister(arena.data_ptr(), arena.numel(), 0) != cudart.cudaError.success: + logger.warning("cudaHostRegister failed; falling back to torch pinned allocations") + return [t.detach().to("cpu").pin_memory() for t in tensors], None + views, offset = [], 0 + for tensor, size in zip(tensors, sizes, strict=True): + nbytes = tensor.numel() * tensor.element_size() + view = arena[offset:offset + nbytes].view(tensor.dtype).view(tensor.shape) + view.copy_(tensor) + views.append(view) + offset += size + return views, arena + + +def _pinned_swap(module: Any, device: torch.device) -> None: + """Move a module's tensors between the GPU and persistent, exactly sized pinned host copies. + + 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 keep a persistent host + copy too and are copied back into it on every park, so mutable buffers stay correct without new allocations. + """ + store = module.__dict__.setdefault("_pinned_host_tensors", {}) + arenas = module.__dict__.setdefault("_pinned_host_arenas", []) + params = dict(module.named_parameters()) + named = [(name, tensor) for name, tensor in list(params.items()) + list(module.named_buffers()) + if tensor is not None] + if device.type == "cpu": + missing = [(name, tensor) for name, tensor in named if tensor.device.type != "cpu" and ( + name not in store or store[name].shape != tensor.shape or store[name].dtype != tensor.dtype)] + if missing: + views, arena = _exact_pinned_views([tensor.detach() for _, tensor in missing]) + store.update({name: view for (name, _), view in zip(missing, views, strict=True)}) + if arena is not None: + arenas.append(arena) + fresh = {name for name, _ in missing} + else: + fresh = set() + for name, tensor in named: + if tensor.device.type == "cpu": + continue + host = store[name] + if name not in params and name not in fresh: + host.copy_(tensor) + tensor.data = host + else: + for _, tensor in named: + if 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 +170,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 +198,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 +207,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 +222,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 +243,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 +362,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 +408,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 +433,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 +489,61 @@ 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 _module_has_dtensor_params(module): + return False + if os.environ.get("FASTVIDEO_H3_PINNED_SWAP", "1") == "1": + _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 = os.environ.get("FASTVIDEO_H3_PARK_MODULES") + 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) -> None: + from fastvideo.pipelines import composed_pipeline_base + + device = composed_pipeline_base.get_local_torch_device() + restored = False + for name in _DENOISE_MODULE_NAMES: + 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() if not self._denoise_stages_ready: self._add_denoise_stages(ref2va=self._ref2va) for name in ( @@ -441,7 +556,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 +596,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 +639,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_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..bdb3f8e546 --- /dev/null +++ b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py @@ -0,0 +1,33 @@ +# 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) + # 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/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_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/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index 02955b73fe..c513ff1239 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -27,19 +27,30 @@ 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 -def _patch_pipeline_construction(monkeypatch, events: list) -> None: + module.device = None + module.moved_to = [] + module.to = to + return module + + +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 +60,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 +114,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 +123,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 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..c4c44a4931 --- /dev/null +++ b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py @@ -0,0 +1,74 @@ +# 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 == [] 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/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index d0712091c6..beb2af9385 100644 --- a/fastvideo/worker/gpu_worker.py +++ b/fastvideo/worker/gpu_worker.py @@ -1,4 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 +import os from typing import Any, cast import torch @@ -21,6 +22,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 +104,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 = os.environ.get("FASTVIDEO_CUDA_MEMORY_CAP_GIB") + if cap_gib: + total = torch.cuda.get_device_properties(self.device).total_memory + torch.cuda.set_per_process_memory_fraction(min(1.0, float(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 +129,8 @@ def init_device(self) -> None: self.distributed_init_method) self.pipeline = build_pipeline(self.fastvideo_args) + if os.environ.get("FASTVIDEO_MEMORY_REPORT") == "1" 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/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..0db7b113a7 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -0,0 +1,609 @@ +"""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]}" + run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}{tag}" + 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..669581d804 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -0,0 +1,116 @@ +"""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}} + env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA"))} + 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, + "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, + "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..8b2ff1b448 --- /dev/null +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -0,0 +1,244 @@ +# 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``). The calibrated weights are +therefore carried over bit for bit; only the activation scale changes, because +FastVideo quantizes activations per call with a unit global scale and the +static ``input_scale`` is dropped. + +``--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" +_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) -> 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), + } + 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", "input_scale")} + 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 + attention = sorted(k[:-len(".weight")] for k in weight_map + if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")])) 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) + 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": len(attention) + len(modelopt) if amax_table else 0, + "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()