diff --git a/docs/contributing/env_vars.md b/docs/contributing/env_vars.md index 4f47ddcb1f..a834261839 100644 --- a/docs/contributing/env_vars.md +++ b/docs/contributing/env_vars.md @@ -192,6 +192,7 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_FA4` | bool | `0` | attention | The FLASH_ATTN backend uses FlashAttention-4 (flash_attn.cute) instead of FA3 or FA2. | | `FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN` | bool | `0` | attention | MiniMax-H3 dense DiT self-attention uses the FlashAttention-4 packed-varlen entry point. This changes the floating-point reduction order, so it is an inference-only opt-in. | | `FASTVIDEO_VSA_SM100A` | bool | `0` | attention | VIDEO_SPARSE_ATTN_H3 sends no-grad tile-64 forwards to the data-center Blackwell (sm_100a) kernel. fastvideo-kernel reads the same variable with the same rule. | +| `FASTVIDEO_VSA_TRITON` | bool | `0` | attention | Force the Triton MiniMax-H3 sparse attention kernel. fastvideo-kernel reads the same variable. | | `FASTVIDEO_NVFP4_FA4` | bool | `0` | attention | FlashAttention-4 quantizes Q and K to NVFP4. An explicit nvfp4_fa4 attention implementation argument takes precedence. | | `FASTVIDEO_DISABLE_ATTENTION_COMPILE` | bool | `1` | attention | Keep attention forward out of torch.compile graphs (torch.compiler.disable). Set it to 0 to let attention constructed under that setting be traced. Setting it explicitly to true also blocks regional compile. | | `FASTVIDEO_MLX_WINDOW` | int | `0` | attention | MLX FastWan windowed attention size in tokens. 0 uses full attention. | @@ -201,6 +202,8 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_VAE_PARALLEL_ENCODE` | bool | `0` | performance | MiniMax-H3 reference-video VAE encode splits its temporal chunks across the sequence-parallel ranks. Same as FastVideoArgs.vae_parallel_encode=True. | | `FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY` | str | unset | performance | Collective that moves chunks in parallel VAE decode: gather (used when unset) or all_gather. | | `FASTVIDEO_MINIMAX_H3_FUSIONS` | str | `""` | performance | MiniMax-H3 inference-only Triton fusions: all, 1, or a comma-separated subset of modulate,qknorm_rope,swiglu. Empty, 0, or none keeps the eager implementation. | +| `FASTVIDEO_H3_VAE_TILE_BATCH` | int | `1` | performance | Spatial tiles per MiniMax-H3 light-VAE decoder call. Values below 1 use one tile. | +| `FASTVIDEO_NVFP4_MM_BACKEND` | str | `auto` | performance | FlashInfer NVFP4 matrix multiplication backend: auto, cutlass, cudnn, trtllm, or b12x. | | `FASTVIDEO_FSDP2_AUTOWRAP` | bool | `0` | performance | FSDP2 shards modules by parameter count instead of the model's shard conditions. Not supported by self-forcing distillation. | | `FASTVIDEO_FSDP2_MIN_PARAMS` | int | `10000000` | performance | Minimum parameter count of a module that FASTVIDEO_FSDP2_AUTOWRAP shards. | | `FASTVIDEO_MLX_COMPILE` | bool | `0` | performance | Compile the MLX DiT forward with mx.compile. | diff --git a/docs/cookbook/minimax-h3.md b/docs/cookbook/minimax-h3.md index 93c8df76c7..b82762f654 100644 --- a/docs/cookbook/minimax-h3.md +++ b/docs/cookbook/minimax-h3.md @@ -11,6 +11,11 @@ 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). +The 42-block pruned checkpoint has an [MLX INT8/INT6 conversion and +eight-forward T2VA command](../getting_started/installation/mlx.md#pruned-eight-forward-checkpoint). +It reads `fastvideo_inference.json` for the trained schedule. The command +uses native 832x480 resolution and all requested frames. +
All model families diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index 359218e583..a2fb02206a 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -46,6 +46,109 @@ is the higher-quality FastH3. Recorded shapes and evidence live in the [support matrix](../../inference/support_matrix.md#apple-silicon-native-runtime). +## Pruned eight-forward checkpoint + +The pruned FastH3 checkpoint has 42 transformer blocks and rank-16 AdaLN. +Its `fastvideo_inference.json` fixes eight denoising forwards, video/audio +shifts of 10/3, and VSA sparsity 0.8. Keep that file beside the transformer +when converting. The converter reads its schedule to build the AdaLN cache. + +```bash +hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ + --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ + --exclude 'text_encoder/*' + +# Optional BF16 encoder fallback: stream the first 50 language layers. +# The last three shards are unused. The packed NVFP4 option is described below. +hf download MiniMaxAI/MiniMax-H3 \ + --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ + --include 'text_encoder/model-0000[1-9]-of-00014.safetensors' \ + --include 'text_encoder/model-0001[0-1]-of-00014.safetensors' \ + --include 'text_encoder/model.safetensors.index.json' \ + --include 'text_encoder/config.json' + +python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \ + --model-root ./FastH3-Pruned-8Step-BF16-ckpt300/transformer \ + --out ./FastH3-Pruned-MLX-vsa \ + --formats "int8 int6" --include-vsa + +python examples/inference/basic/mlx_fasth3.py \ + --model-root ./FastH3-Pruned-8Step-BF16-ckpt300 \ + --mlx-checkpoint ./FastH3-Pruned-MLX-vsa/int8 \ + --prompt "(S1) A potter asks [English] Is the rim ready?" \ + --height 480 --width 832 --num-frames 243 --steps 8 \ + --vsa --vsa-sparsity 0.8 --vsa-tile-size 64 \ + --output-path ./outputs/fasth3_pruned_int8_480p.mp4 +``` + +At 24 fps, 124 frames is the legal H3 count for a roughly five-second clip. +Use `--num-frames 124` and a separate output path for that run. The `--fast` +and `--fast-spatial` options change the workload and are not part of the +native-resolution benchmark. A 36 GB Mac may need INT6 and phased loading; +measure memory before claiming all-resident operation. + +### Packed encoder and resident loading + +The experimental MLX conditioner can read the released FastVideo NVFP4 +text encoder directly, using native `nvfp4` matrix multiplication. It keeps +the packed weights and BF16 embedding table in memory, with FP32 +activations. CUDA uses quantized activations, so the two encoders are not +bit-exact. Validate generated video and audio before publishing a timing. +MLX 0.32.2 supports the required operator on Apple Silicon. + +Pass the packed encoder directory as `conditioner_dir`; `conditioner_mode="auto"` +selects it from `config.json`. The BF16 fallback continues to stream layers. +To request all-resident generation through the Python API: + +```python +from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline + +pipeline = MiniMaxH3MLXPipeline( + model_root="./FastH3-Pruned-8Step-BF16-ckpt300", + mlx_dit_checkpoint="./FastH3-Pruned-MLX-vsa/int6", + conditioner_dir="./FastH3-NVFP4-encoder", + conditioner_mode="nvfp4", + resident=True, + vae_dtype="fp16", +) +try: + pipeline.prepare_resident() # Load encoder, DiT, video VAE and audio VAE. + result = pipeline.generate( + "(S1) A potter asks [English] Is the rim ready?", + output_path="./outputs/fasth3_pruned_resident.mp4", + height=480, width=832, num_frames=243, num_steps=8, + vsa=True, vsa_sparsity=0.8, vsa_tile_size=64, + ) +finally: + pipeline.close() +``` + +Resident placement requires space for activations as well as all four +components. On a 36 GiB Mac, try INT6 first and measure peak allocation. +If loading or inference runs out of memory, use phased loading by leaving +`resident=False`. Changing placement does not change frames or resolution. + +### Metal wired memory + +MLX's allocation limit and wired-memory limit are separate. H3's optional +`metal_wired_limit_gib` calls `mx.set_wired_limit` so selected Metal allocations +stay in physical memory. It does not increase available RAM. Explicit requests +fail visibly if the installed MLX build cannot apply them. + +Inspect the device's recommended working set before choosing a limit: + +```python +import mlx.core as mx + +print(mx.device_info()) +``` + +For the tested 36 GiB M4 Max, phased generation can request +`metal_wired_limit_gib=27` in `MiniMaxH3MLXPipeline`. Leave room for macOS and +other applications. All-resident generation also needs room for the encoder, +DiT, both decoders and peak activations; wiring cannot make an oversized stack +fit. An omitted wired limit preserves MLX's existing wiring setting. + ## Hardware - FastMetal 1.3B and 5B: 16 GB unified memory and up diff --git a/docs/getting_started/installation/spark.md b/docs/getting_started/installation/spark.md index 33888a853a..570cf71f2c 100644 --- a/docs/getting_started/installation/spark.md +++ b/docs/getting_started/installation/spark.md @@ -144,6 +144,9 @@ for which models are practical on the GB10, what makes them faster, and what won't help on this hardware (and why) — so you don't spend a night tuning knobs that can't move here. +For the eight-forward FastH3 V2 NVFP4 stack with a trimmed encoder and light +VAE, use the [one-Spark resident recipe](spark_performance.md#fasth3-v2-nvfp4-on-one-spark). + Two Sparks with QSFP cables: [Pair two NVIDIA DGX Sparks](spark_pair.md) for one FastH3 clip across both GPUs (`sp_size=2` over Ray). Copy-paste commands for one or two Sparks also live on the diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index e0f5cb4d13..581f10f9d8 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -161,14 +161,16 @@ is power-cycled. To avoid it: on: "CPU" offload uses the same unified RAM. Multi-GPU FSDP sharding remains available because it partitions weights without parking them in a separate host pool. -- **MiniMax H3 / FastH3** still needs deferred loading on one GB10. The Qwen3-VL - conditioner is tens of gigabytes of BF16. If the DiT and VAEs load while that - encoder is still resident, the process is a typical `earlyoom` kill (Python is - preferred). On unified memory, `lazy_module_load` auto-enables and owns that +- **Older MiniMax H3 / FastH3 bf16 weights** need deferred loading on one GB10. + The full Qwen3-VL conditioner is tens of gigabytes of BF16. If the DiT and + VAEs load while that encoder is still resident, the process can be killed by + `earlyoom`. On unified memory, `lazy_module_load` auto-enables and owns that split (encoder, then DiT, then VAE; DiT can drop before decode). Sequential - load is the H3-only fallback when lazy is off; do not pass - `--no-lazy-module-load` here. Geometry scalars come from checkpoint - `config.json`, not live weights. See [Offloading](../../inference/offloading.md). + load is the H3-only fallback when lazy is off. Keep deferred loading for + those older checkpoints. The trimmed NVFP4 encoder and light VAE in the + [V2 resident recipe](#fasth3-v2-nvfp4-on-one-spark) are a different memory + profile. Geometry scalars come from checkpoint `config.json`, not live + weights. See [Offloading](../../inference/offloading.md). - **FastH3 TAEH3** (`--video-decode-backend taeh3`) is an opt-in preview decoder. T2VA never materializes the 9.7 GiB video VAE (DiT still loads after Qwen via sequential start). On this box, alpine 768×1344×124 decoded in **2.4 s** versus @@ -207,6 +209,147 @@ A few things that surprise people on this box (beyond the memory notes above): `Released MiniMax-H3 text encoder after conditioning` before `Loading MiniMax-H3 denoise modules`). +## FastH3 V2 NVFP4 on one Spark + +This recipe uses the full V2 eight-forward transformer, the 50-layer NVFP4 +Qwen3-VL encoder, and the light H3 video VAE. Its configuration keeps all +three resident on one GB10. This stack fits in the Spark's unified memory; +benchmark your installed runtime and review the clips before publishing a +speed claim. The earlier bf16 H3 memory guidance above concerns a larger checkpoint. + +Install FastVideo from a checkout that includes the ModelOpt converter and +FlashInfer FP4 support, following [the Spark install guide](spark.md). Sign in +to Hugging Face with access to the FastVideo model repositories. Download the +V2 scheduler and audio components, the compact encoder and VAE from the pruned +repo, and the ModelOpt V2 transformer. The pruned model's encoder and VAE are +the same components used by V2. + +```bash +SPARK_STACK=./FastH3-V2-Spark-NVFP4 +V2_FP4_SRC=./FastH3-V2-ModelOpt-NVFP4 + +hf download FastVideo/FastVideo-FastH3-8-Step-V2 \ + --local-dir "$SPARK_STACK" \ + --exclude 'transformer/*' --exclude 'text_encoder/*' --exclude 'vae/*' +hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ + --local-dir "$SPARK_STACK" \ + --include 'text_encoder/*' --include 'vae/*' +hf download FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4 \ + --local-dir "$V2_FP4_SRC" --include 'transformer/*' + +nice -n 19 python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \ + --src "$V2_FP4_SRC/transformer" --dst "$SPARK_STACK/transformer" \ + --quantize-attention --quantize-gate + +test -f "$SPARK_STACK/transformer/nvfp4_weights.safetensors" +test -f "$SPARK_STACK/text_encoder/config.json" +test -f "$SPARK_STACK/vae/config.json" +test -f "$SPARK_STACK/fastvideo_inference.json" +python -m json.tool "$SPARK_STACK/fastvideo_inference.json" >/dev/null +``` + +The converter probes each packed linear through FlashInfer `mm_fp4`. If that +probe fails on `sm_121`, convert the transformer on another Blackwell GPU and +copy the resulting `transformer/` directory to the Spark. Do not omit +`fastvideo_inference.json`: it supplies V2's trained denoising ladder. The +recipe's `num_inference_steps: 9` means nine sigma points and eight DiT +forwards. + +The converter preserves ModelOpt's calibrated `input_scale` as the reciprocal +`_nvfp4_input_global_sf`. Reconvert older exports that discarded this scale: +unit activation scaling clips inputs above 2688. An explicit `--act-amax` +table overrides the source calibration. + +On GB10 with FlashInfer 0.6.18, FastVideo fences activation quantization before +releasing its padded input. Without this completion fence, identical H3 +requests produced different DiT latents and occasionally corrupt video. +The fence applies to `sm_121`; other architectures retain asynchronous execution. + +Run `examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml` from the +repository root. It uses 832x480, 243 frames, VSA sparsity 0.8 with +64-token tiles, and the light H3 VAE through the `h3-vae` decode backend. +It does not use frame dropping or spatial upscaling. + +```bash +FASTVIDEO_MINIMAX_H3_FUSIONS=all \ +FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +FASTVIDEO_H3_VAE_TILE_BATCH=1 \ +FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 \ +FASTVIDEO_STAGE_LOGGING=1 \ +nice -n 19 fastvideo generate \ + --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml +``` + +For a roughly five-second clip, set `--request.sampling.num_frames 124` and +write to a separate output path. H3 permits frame counts of `17n+5`; 124 is +the closest legal count above five seconds at 24 fps. For the secondary +10-second setting, set `--request.sampling.width 1344` and +`--request.sampling.height 768`, keeping 243 frames. Use the two prompts in +`handoff_spark_mac/benchmark_prompts.json` from the local release handoff. +The benchmark script runs one warmup and at least two timed generations for +each prompt in one process. It saves the MP4s and prints the wall time, stage +times, peak memory, and median. Set the Spark environment before running it: + +```bash +export FASTVIDEO_MINIMAX_H3_FUSIONS=all +export FASTVIDEO_NVFP4_MM_BACKEND=cutlass FASTVIDEO_H3_VAE_TILE_BATCH=1 +export FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 +export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 +nice -n 19 python examples/inference/basic/benchmark_fasth3_spark_nvfp4.py \ + --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml \ + --prompts /path/to/fasth3-local-release/handoff_spark_mac/benchmark_prompts.json \ + --output-dir outputs/fasth3_spark_v2_nvfp4/benchmark-243 --frames 243 + +# Repeat with --frames 124 and a different output directory for the five-second check. +``` + +Record the exact command and commit with the measurements. Review every clip's +video and audio before publishing a quality or speed claim. + +After the V2 baseline works, sweep `FASTVIDEO_H3_VAE_TILE_BATCH` and +`FASTVIDEO_NVFP4_MM_BACKEND` on the same prompts. Compare the optional AdaLN +table and VAE compile only with the same frame count, schedule, and VSA +sparsity. The V2 converter packs VSA gates, so its `h3_dit_vsa` profile must +match the recipe. A later pruned NVFP4 transformer uses the separate +`h3_dit_ffn` profile, with attention and VSA gates left dense. + +### Native FastH3 release measurements + +Measured on October 4, 2026, with the trained eight-forward ladder, VSA 0.8, +832x480 native video, the 50-layer NVFP4 encoder and light video/audio VAEs +resident. Each cell is the median of two timed calls after one warmup, in +seconds, for `latency-ceramics-005` / `latency-harbor-005`, seed 2026. + +| Model | Frames | One Spark | Two Sparks, SP2 | +|---|---:|---:|---:| +| Pruned ckpt300, FFN NVFP4 | 124 | 134.361 / 134.675 | **78.245 / 78.272** | +| Pruned ckpt300, FFN NVFP4 | 243 | 284.188 / 280.593 | **165.987 / 163.001** | +| Full V2, NVFP4 FFN/attention/gates | 124 | 142.460 / 140.425 | **88.356 / 86.101** | +| Full V2, NVFP4 FFN/attention/gates | 243 | 308.204 / 306.141 | **179.876 / 180.213** | + +The base integrates upstream main `0cc41a22` with experimental NVFP4 support. +One-Spark tested commits are `6d6b57fe` (pruned 124), `3f24557a` (pruned 243, +with `CUDA_LAUNCH_BLOCKING=1`), `f5126f78` (V2 124) and `6e9d7a0b` (V2 243). +The final pair uses `715d4a5f`, with matching actual-worker code fingerprints, +CUTLASS FP4 GEMMs and Triton VSA. Light-VAE tile batch is 8 on one Spark and 1 +on the pair; pair batch 8 was slower at 124 frames (79.993 / 80.022 s pruned). +All offload/deferred-loading and compile options are disabled. The pair uses +QSFP RoCE, SP2/TP1 and parallel VAE gathering. See the pair configs +`basic_fasth3_spark_pair_pruned_nvfp4.yaml` and +`basic_fasth3_spark_pair_v2_nvfp4.yaml` beside the benchmark script. + +Every final pair warmup and repeat has correct dimensions/frame count, coherent +sampled frames and identical full decoded-video hashes within its prompt. +The V2 one-Spark 124-frame harbor warmup differs from the timed clips but remains +coherent. These checks establish repeat reliability for the tested recipes; +BF16 reference parity, speech accuracy and lip sync need separate review. + +The [older H3 local blog](https://haoailab.com/blogs/fasth3-local/) reports +243 s on one Spark and 209 s on two Sparks at 124 frames. It uses the four-step +Preview checkpoint and full VAE, so the old and new values are context, not a +matched optimization comparison. It has no matching 243-frame baseline. + ## Reproduce these numbers Two scripts under `examples/inference/optimizations/` reproduce the claims on diff --git a/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml new file mode 100644 index 0000000000..32852ad514 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml @@ -0,0 +1,63 @@ +# Start Ray on both Sparks and source spark_pair_env.sh first; see spark_pair.md. +# FastH3 pruned ckpt300 eight-forward video+audio on two DGX Sparks over QSFP RoCE. +# Download the complete checkpoint, including its trained schedule, as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer uses NVFP4 FFN weights with bf16 attention +# and VSA gates (layer_profile: h3_dit_ffn). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml +generator: + model_path: FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 + engine: + num_gpus: 2 + execution_backend: ray + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_ffn + parallelism: + tp_size: 1 + sp_size: 2 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: true + vae_parallel_decode_strategy: gather + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pair_pruned_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml new file mode 100644 index 0000000000..89bef8ed8a --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml @@ -0,0 +1,63 @@ +# Start Ray on both Sparks and source spark_pair_env.sh first; see spark_pair.md. +# FastH3 V2 eight-forward video+audio on two DGX Sparks over QSFP RoCE. +# Assemble ./FastH3-V2-Spark-NVFP4 as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer must include NVFP4 attention +# and VSA gates (layer_profile: h3_dit_vsa). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml +generator: + model_path: ./FastH3-V2-Spark-NVFP4 + engine: + num_gpus: 2 + execution_backend: ray + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_vsa + parallelism: + tp_size: 1 + sp_size: 2 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: true + vae_parallel_decode_strategy: gather + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pair_v2_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml new file mode 100644 index 0000000000..a5d1af4c47 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml @@ -0,0 +1,60 @@ +# FastH3 pruned ckpt300 eight-forward video+audio on one DGX Spark. +# Download the complete checkpoint, including its trained schedule, as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer uses NVFP4 FFN weights with bf16 attention +# and VSA gates (layer_profile: h3_dit_ffn). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml +generator: + model_path: FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_ffn + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pruned_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml new file mode 100644 index 0000000000..c3986fabd0 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml @@ -0,0 +1,60 @@ +# FastH3 V2 eight-forward video+audio on one DGX Spark. +# Assemble ./FastH3-V2-Spark-NVFP4 as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer must include NVFP4 attention +# and VSA gates (layer_profile: h3_dit_vsa). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml +generator: + model_path: ./FastH3-V2-Spark-NVFP4 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_vsa + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_v2_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py new file mode 100644 index 0000000000..f8e3ee7fcd --- /dev/null +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -0,0 +1,144 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Time a resident FastH3 eight-forward Spark recipe with the release prompts. + +One process loads the model, then each prompt gets one excluded warmup and at +least two timed calls. Every call writes a video. This script does not alter +the V2 schedule, VSA sparsity, or video resolution. +""" + +from __future__ import annotations + +import argparse +import json +import statistics +import time +from copy import deepcopy +from pathlib import Path + +from fastvideo import VideoGenerator +from fastvideo.api.parser import load_raw_config, parse_config +from fastvideo.api.schema import RunConfig + +PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") + + +def _stage_metrics(result: object) -> dict[str, dict]: + logging_info = getattr(result, "logging_info", None) + stages = getattr(logging_info, "stages", None) + if isinstance(logging_info, dict): + stages = logging_info.get("stages", stages) + if not isinstance(stages, dict): + return {} + return {name: metrics for name, metrics in stages.items() if isinstance(metrics, dict)} + + +def _stage_seconds(metrics: dict[str, dict]) -> dict[str, float]: + return {name: float(stage["execution_time"]) for name, stage in metrics.items() + if stage.get("execution_time") is not None} + + +def _peak_mb(metrics: dict[str, dict], key: str) -> float | None: + values = [float(stage[key]) for stage in metrics.values() if stage.get(key) is not None] + return max(values) if values else None + + +def _stage_total(stages: dict[str, float], fragment: str) -> float | None: + matches = [seconds for name, seconds in stages.items() if fragment in name.lower()] + return sum(matches) if matches else None + + +def _request(base: RunConfig, prompt: str, frames: int, width: int, height: int, + output: Path): + request = deepcopy(base.request) + request.prompt = prompt + request.inputs.prompt_path = None + request.sampling.num_frames = frames + request.sampling.width = width + request.sampling.height = height + request.output.output_path = str(output) + request.output.save_video = True + request.output.return_frames = False + return request + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", type=Path, required=True) + parser.add_argument("--prompts", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--model-path", type=Path) + parser.add_argument("--frames", type=int, default=243) + parser.add_argument("--width", type=int, default=832) + parser.add_argument("--height", type=int, default=480) + parser.add_argument("--repeats", type=int, default=2) + args = parser.parse_args() + + if args.frames not in (124, 243): + parser.error("use 124 frames for roughly five seconds or 243 for the ten-second headline") + if args.repeats < 2: + parser.error("the release protocol requires at least two timed calls") + + config = parse_config(RunConfig, load_raw_config(args.config)) + if args.model_path: + config.generator.model_path = str(args.model_path) + if config.request.sampling.num_inference_steps != 9: + parser.error("the V2 contract requires nine sigma points for eight DiT forwards") + if config.generator.engine.offload.lazy_module_load is not False: + parser.error("the resident recipe requires lazy_module_load: false") + contract = Path(config.generator.model_path) / "fastvideo_inference.json" + if not contract.is_file(): + parser.error(f"missing trained V2 schedule: {contract}") + inference = json.loads(contract.read_text()) + if inference.get("num_inference_steps") != 9 or inference.get("transformer_forwards") != 8: + parser.error("the checkpoint is not the trained V2 eight-forward schedule") + + prompts = json.loads(args.prompts.read_text()) + if any(prompt_id not in prompts for prompt_id in PROMPT_IDS): + parser.error(f"prompt JSON must contain {', '.join(PROMPT_IDS)}") + args.output_dir.mkdir(parents=True, exist_ok=True) + + generator = VideoGenerator.from_config(config.generator) + try: + for prompt_id in PROMPT_IDS: + times = [] + for index in range(args.repeats + 1): + warmup = index == 0 + label = "warmup" if warmup else f"run-{index:02d}" + requested_path = args.output_dir / f"{prompt_id}-{args.width}x{args.height}-{args.frames}-{label}.mp4" + request = _request(config, prompts[prompt_id], args.frames, args.width, args.height, + requested_path) + started = time.perf_counter() + result = generator.generate(request) + wall = time.perf_counter() - started + output = Path(result.video_path) if result.video_path else requested_path + if not output.is_file(): + raise RuntimeError(f"generation returned without an MP4: {output}") + metrics = _stage_metrics(result) + stages = _stage_seconds(metrics) + row = { + "prompt_id": prompt_id, + "warmup": warmup, + "frames": args.frames, + "width": args.width, + "height": args.height, + "e2e_seconds": round(wall, 3), + "denoise_seconds": _stage_total(stages, "denois"), + "decode_seconds": _stage_total(stages, "decod"), + "peak_memory_mb": _peak_mb(metrics, "peak_allocated_mb"), + "peak_reserved_mb": _peak_mb(metrics, "peak_reserved_mb"), + "result_peak_memory_mb": result.peak_memory_mb, + "stages": stages, + "mp4": str(output), + } + print(json.dumps(row, sort_keys=True), flush=True) + if not warmup: + times.append(wall) + print(json.dumps({"prompt_id": prompt_id, "timed_runs": len(times), + "median_e2e_seconds": round(statistics.median(times), 3)}, + sort_keys=True), flush=True) + finally: + generator.shutdown() + + +if __name__ == "__main__": + main() diff --git a/fastvideo/envs.py b/fastvideo/envs.py index 8962a8e703..7c0aceb73c 100644 --- a/fastvideo/envs.py +++ b/fastvideo/envs.py @@ -331,6 +331,10 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="attention", doc="VIDEO_SPARSE_ATTN_H3 sends no-grad tile-64 forwards to the data-center Blackwell (sm_100a) kernel. " "fastvideo-kernel reads the same variable with the same rule.") +FASTVIDEO_VSA_TRITON = EnvBool( + False, + category="attention", + doc="Force the Triton MiniMax-H3 sparse attention kernel. fastvideo-kernel reads the same variable.") FASTVIDEO_NVFP4_FA4 = EnvBool( False, category="attention", @@ -378,6 +382,12 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="performance", doc="MiniMax-H3 inference-only Triton fusions: all, 1, or a comma-separated subset of " "modulate,qknorm_rope,swiglu. Empty, 0, or none keeps the eager implementation.") +FASTVIDEO_H3_VAE_TILE_BATCH = EnvInt( + 1, category="performance", doc="Spatial tiles per MiniMax-H3 light-VAE decoder call. Values below 1 use one tile.") +FASTVIDEO_NVFP4_MM_BACKEND = EnvStr( + "auto", + category="performance", + doc="FlashInfer NVFP4 matrix multiplication backend: auto, cutlass, cudnn, trtllm, or b12x.") FASTVIDEO_FSDP2_AUTOWRAP = EnvBool(False, category="performance", doc="FSDP2 shards modules by parameter count instead of the model's shard " diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 403af3b5ac..5735048d10 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -28,6 +28,7 @@ import torch.nn.functional as F from torch.nn.parameter import Parameter +import fastvideo.envs as envs from fastvideo.layers.quantization.base_config import ( QuantizationConfig, QuantizeMethodBase, @@ -191,8 +192,27 @@ def _nvfp4_quantize_op( sf_layout: int, do_shuffle: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: + spark = torch.cuda.get_device_capability(x.device) == (12, 1) + if spark and torch.cuda.is_current_stream_capturing(): + raise RuntimeError("NVFP4 activation quantization on DGX Spark requires a completion fence; " + "disable CUDA graph capture.") SfLayout, _, nvfp4_quantize = _require_flashinfer() - return nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) + if spark: + # FlashInfer's PDL kernel reads the global scale before its + # dependency wait. Fresh dynamic scales require normal ordering. + quantized, scales = nvfp4_quantize(x, + global_sf, + sfLayout=SfLayout(sf_layout), + do_shuffle=do_shuffle, + enable_pdl=False) + else: + quantized, scales = nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) + if spark: + # With FlashInfer 0.6.18 on GB10, queued activation quantization + # plus GEMM can diverge. Completing quantization while its padded + # input is alive prevents the observed intermittent corruption. + torch.cuda.current_stream(x.device).synchronize() + return quantized, scales @_nvfp4_quantize_op.register_fake def _nvfp4_quantize_op_fake( @@ -347,7 +367,7 @@ def _mm_fp4_backend() -> str: ``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") + return envs.FASTVIDEO_NVFP4_MM_BACKEND.get() def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: @@ -380,6 +400,8 @@ def _load_amax_table(path: str) -> dict[str, float]: class NVFP4QuantizeMethod(QuantizeMethodBase): + _static_sf: torch.Tensor | None + def __init__(self, layer_prefix: str = ""): super().__init__() self.weight_fp4 = None @@ -418,8 +440,7 @@ def _static_activation_global_sf(self) -> torch.Tensor | None: 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") + 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: diff --git a/fastvideo/mlx_runtime/fastwan.py b/fastvideo/mlx_runtime/fastwan.py index d6a802b2b1..1cf8844ebc 100644 --- a/fastvideo/mlx_runtime/fastwan.py +++ b/fastvideo/mlx_runtime/fastwan.py @@ -90,6 +90,7 @@ class QuantizedMatrix: biases: mx.array | None spec: MLXQuantizationSpec dequantized_dtype: mx.Dtype + global_scale: float = 1.0 def fastwan_shape( @@ -281,14 +282,24 @@ def ensure_quantization_supported(spec: MLXQuantizationSpec | None) -> None: f"(int8 is currently the most reliable quality/memory target).") -def quantize_matrix(weight, spec: MLXQuantizationSpec | None): +def quantize_matrix(weight, spec: MLXQuantizationSpec | None, *, use_nvfp4_global_scale: bool = False): if spec is None: return weight import mlx.core as mx if len(weight.shape) < 2: return weight - q = mx.quantize(weight, group_size=spec.group_size, bits=spec.bits, mode=spec.mode) + global_scale = 1.0 + quantization_input = weight + if spec.mode == "nvfp4" and use_nvfp4_global_scale: + # E4M3 block scales cannot represent typical small model weights + # directly. Normalize into E2M1's max 6 times E4M3's max 448. + maximum = float(mx.max(mx.abs(weight)).item()) + global_scale = maximum / (6.0 * 448.0) if maximum > 0 else 1.0 + # Normalize explicitly so this also works with older MLX operators + # without the global_scale keyword. Matmul restores this multiplier. + quantization_input = weight.astype(mx.float32) / global_scale + q = mx.quantize(quantization_input, group_size=spec.group_size, bits=spec.bits, mode=spec.mode) biases = q[2] if len(q) == 3 else None eval_args = [q[0], q[1]] if biases is not None: @@ -300,6 +311,7 @@ def quantize_matrix(weight, spec: MLXQuantizationSpec | None): biases=biases, spec=spec, dequantized_dtype=weight.dtype, + global_scale=global_scale, ) @@ -371,7 +383,7 @@ def _quantized_linear(x, weight: QuantizedMatrix, *, use_affine_dq_gemm: bool = _dq_gemm_logged = True logger.info("affine dequant+GEMM engaged (rows=%d, floor=%d, bits=%s)", rows, min_m, spec.bits) return y - return mx.quantized_matmul( + result = mx.quantized_matmul( x, weight.weight, weight.scales, @@ -381,6 +393,9 @@ def _quantized_linear(x, weight: QuantizedMatrix, *, use_affine_dq_gemm: bool = bits=spec.bits, mode=spec.mode, ).astype(x.dtype) + if weight.global_scale != 1.0: + result = (result.astype(mx.float32) * weight.global_scale).astype(x.dtype) + return result def linear(x, weight, bias=None, *, use_affine_dq_gemm: bool = False): diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 58b677bd80..37ff66b215 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -62,7 +62,7 @@ QuantizedMatrix, ensure_quantization_supported, linear as _shared_linear, - quantize_matrix, + quantize_matrix as _shared_quantize_matrix, silu, timestep_embedding, weight_dtype, @@ -87,6 +87,11 @@ def linear(x, weight, bias=None): return _shared_linear(x, weight, bias, use_affine_dq_gemm=True) +def quantize_matrix(weight, spec: MLXQuantizationSpec | None): + """Use a global NVFP4 scale for H3's small transformer weights.""" + return _shared_quantize_matrix(weight, spec, use_nvfp4_global_scale=True) + + # --------------------------------------------------------------------------- # Constants (mirrors fastvideo/pipelines/basic/minimax_h3/packing.py) # --------------------------------------------------------------------------- @@ -766,12 +771,12 @@ def _feed_forward(weights: dict[str, Any], x): return linear(value * silu(gate), weights["ff.net.2.weight"]) -def _adaln_tables(weights: dict[str, Any], temb): +def _adaln_tables(weights: dict[str, Any], temb, *, apply_silu: bool = True): """Six (n_t * 3, hidden) modulation tables from (n_t, time_embed_dim).""" import mlx.core as mx projected = linear( - silu(temb).astype(weight_dtype(weights["adaln_proj.linear.weight"])), + (silu(temb) if apply_silu else temb).astype(weight_dtype(weights["adaln_proj.linear.weight"])), weights["adaln_proj.linear.weight"], weights["adaln_proj.linear.bias"], ) @@ -914,6 +919,7 @@ def __init__( self.qk_norm_eps = float(config["qk_norm_eps"]) self.final_norm_eps = float(config["final_norm_eps"]) self.patch_dim = self.in_channels * math.prod(self.patch_size) + self.adaln_rank = config.get("adaln_rank") self._adaln_cache: MiniMaxH3StepCache | None = None self.vsa_config = MiniMaxH3VSAConfig() self._vsa_geometry: MiniMaxH3VSAGeometry | None = None @@ -932,11 +938,15 @@ def compute_temb(self, timesteps): self.weights["time_embedder.linear_1.weight"], self.weights["time_embedder.linear_1.bias"], ) - return linear( + temb = linear( silu(temb), self.weights["time_embedder.linear_2.weight"], self.weights["time_embedder.linear_2.bias"], ) + if self.adaln_rank is not None: + temb = linear( + silu(temb).astype(weight_dtype(self.weights["adaln_basis.weight"])), self.weights["adaln_basis.weight"]) + return temb def refine_text(self, text_rows): hidden = linear( @@ -969,9 +979,10 @@ def precompute_adaln(self, timesteps: np.ndarray, *, drop_weights: bool = True) timesteps = np.unique(np.asarray(timesteps, dtype=np.float32)) temb = self.compute_temb(mx.array(timesteps)) - block_tables = [_adaln_tables(block, temb) for block in self.blocks] + block_tables = [_adaln_tables(block, temb, apply_silu=self.adaln_rank is None) for block in self.blocks] shift_scale = linear( - silu(temb).astype(weight_dtype(self.weights["norm_out.linear.weight"])), + (silu(temb) if self.adaln_rank is None else temb).astype( + weight_dtype(self.weights["norm_out.linear.weight"])), self.weights["norm_out.linear.weight"], self.weights["norm_out.linear.bias"], ) @@ -1107,7 +1118,7 @@ def forward( adaln_indices = (timestep_indices * MINIMAX_H3_MODALITY_NUM + token_tags).astype(mx.int32) for block_index, block in enumerate(self.blocks): - tables = _adaln_tables(block, temb) + tables = _adaln_tables(block, temb, apply_silu=self.adaln_rank is None) packed = _transformer_block( block, packed, @@ -1124,7 +1135,8 @@ def forward( mx.eval(packed) # per-block sync: see forward_with_cache note shift_scale = linear( - silu(temb).astype(weight_dtype(self.weights["norm_out.linear.weight"])), + (silu(temb) if self.adaln_rank is None else temb).astype( + weight_dtype(self.weights["norm_out.linear.weight"])), self.weights["norm_out.linear.weight"], self.weights["norm_out.linear.bias"], ) @@ -1353,10 +1365,10 @@ def assign(key: str, value) -> None: for shard in _safetensors_shards(transformer_path): shard_arrays = mx.load(str(shard)) for key, source in shard_arrays.items(): - if not key.startswith("time_embedder."): + if not (key.startswith("time_embedder.") or key == "adaln_basis.weight"): continue keep_fp32 = key.split(".", 1)[0] in FP32_MODULE_PREFIXES - target_dtype = mx.float32 if keep_fp32 else cast_dtype + target_dtype = mx.float32 if keep_fp32 else (mx.float16 if key == "adaln_basis.weight" else cast_dtype) assign(key, _load_array(source, target_dtype)) del shard_arrays required_time_keys = { @@ -1373,15 +1385,21 @@ def assign(key: str, value) -> None: weight_dtype(weights["time_embedder.linear_1.weight"])) temb = linear(t_freq, weights["time_embedder.linear_1.weight"], weights["time_embedder.linear_1.bias"]) temb = linear(silu(temb), weights["time_embedder.linear_2.weight"], weights["time_embedder.linear_2.bias"]) + if config.get("adaln_rank") is not None: + if "adaln_basis.weight" not in weights: + raise KeyError("Rank-reduced AdaLN checkpoint is missing adaln_basis.weight") + temb = linear(silu(temb).astype(weight_dtype(weights["adaln_basis.weight"])), weights["adaln_basis.weight"]) mx.eval(temb) cached_block_tables = [None] * num_blocks for shard in _safetensors_shards(transformer_path): shard_arrays = mx.load(str(shard)) for key, source in shard_arrays.items(): + if key.endswith(".weight_scale"): + continue # paired with its FP8 weight if _is_ignored_dense_key(key, include_vsa=include_vsa): continue - if temb is not None and key.startswith("time_embedder."): + if temb is not None and (key.startswith("time_embedder.") or key == "adaln_basis.weight"): continue if key.startswith("transformer_blocks."): index = int(key.split(".")[1]) @@ -1390,15 +1408,27 @@ def assign(key: str, value) -> None: if key.startswith("rope."): continue # non-persistent analytic buffer, rebuilt on the fly keep_fp32 = key.split(".", 1)[0] in FP32_MODULE_PREFIXES - target_dtype = mx.float32 if keep_fp32 else cast_dtype - array = _load_array(source, target_dtype) + factorized_adaln = config.get("adaln_rank") is not None and (".adaln_proj." in key or key.startswith( + ("norm_out.linear.", "adaln_basis."))) + target_dtype = mx.float32 if keep_fp32 else (mx.float16 if factorized_adaln else cast_dtype) + if source.dtype == mx.uint8 and key.endswith(".weight"): + scale_key = key + "_scale" + if scale_key not in shard_arrays: + raise KeyError(f"FP8 weight {key} needs {scale_key} in the same safetensors shard") + scale = shard_arrays[scale_key].astype(mx.float32) + if scale.size != source.shape[0]: + raise ValueError(f"FP8 scale for {key} has {scale.size} entries, expected {source.shape[0]}") + array = (mx.from_fp8(source, dtype=mx.float16) * scale.reshape(-1, 1)).astype(target_dtype) + mx.eval(array) + else: + array = _load_array(source, target_dtype) if temb is not None and ".adaln_proj.linear." in key: _, index_str, sub = key.split(".", 2) index = int(index_str) block_pending = pending_adaln.setdefault(index, {}) block_pending[sub] = array if {"adaln_proj.linear.weight", "adaln_proj.linear.bias"} <= block_pending.keys(): - tables = _adaln_tables(block_pending, temb) + tables = _adaln_tables(block_pending, temb, apply_silu=config.get("adaln_rank") is None) mx.eval(tables) assert cached_block_tables is not None cached_block_tables[index] = tables @@ -1428,7 +1458,8 @@ def assign(key: str, value) -> None: if missing_cache_blocks: raise KeyError(f"Missing AdaLN cache tables for blocks {missing_cache_blocks}") shift_scale = linear( - silu(temb).astype(weight_dtype(weights["norm_out.linear.weight"])), + (silu(temb) if config.get("adaln_rank") is None else temb).astype( + weight_dtype(weights["norm_out.linear.weight"])), weights["norm_out.linear.weight"], weights["norm_out.linear.bias"], ) @@ -1467,6 +1498,7 @@ def assign(key: str, value) -> None: H3_FORMAT_VERSION = 1 +H3_SCALED_NVFP4_FORMAT_VERSION = 2 H3_WEIGHTS_FILENAME = "mlx_h3_dit.safetensors" H3_MANIFEST_FILENAME = "mlx_h3_dit.json" @@ -1536,6 +1568,8 @@ def save_mlx_h3_checkpoint(dit: MLXMiniMaxH3DiT, checkpoint_dir: str | Path) -> "dequantized_dtype": _dtype_name(value.dequantized_dtype), "has_biases": value.biases is not None, } + if value.global_scale != 1.0: + quantized[key]["global_scale"] = value.global_scale else: arrays[key] = value @@ -1554,17 +1588,25 @@ def save_mlx_h3_checkpoint(dit: MLXMiniMaxH3DiT, checkpoint_dir: str | Path) -> arrays["__adaln_cache.norm_out_scale"] = cache.norm_out_scale manifest = { - "format_version": H3_FORMAT_VERSION, - "config": dit.config, - "num_blocks": len(dit.blocks), - "num_refiner_blocks": len(dit.refiner), - "quantization": None if spec is None else { + "format_version": + (H3_SCALED_NVFP4_FORMAT_VERSION if any("global_scale" in info + for info in quantized.values()) else H3_FORMAT_VERSION), + "config": + dit.config, + "num_blocks": + len(dit.blocks), + "num_refiner_blocks": + len(dit.refiner), + "quantization": + None if spec is None else { "mode": spec.mode, "bits": spec.bits, "group_size": spec.group_size, }, - "quantized_keys": quantized, - "adaln_cache": cache_manifest, + "quantized_keys": + quantized, + "adaln_cache": + cache_manifest, "vsa": { "capable": bool(dit.vsa_capable), @@ -1603,9 +1645,9 @@ def load_mlx_h3_checkpoint(checkpoint_dir: str | Path) -> MLXMiniMaxH3DiT: manifest = json.loads(manifest_path.read_text()) version = manifest.get("format_version") - if version != H3_FORMAT_VERSION: + if version not in (H3_FORMAT_VERSION, H3_SCALED_NVFP4_FORMAT_VERSION): raise ValueError(f"MLX H3 checkpoint {checkpoint_dir} has format_version={version}; " - f"this build reads version {H3_FORMAT_VERSION}. Re-export the checkpoint.") + f"this build reads versions 1 and 2. Re-export the checkpoint.") spec = None if manifest["quantization"] is not None: @@ -1620,12 +1662,16 @@ def rebuild(key: str): return arrays[key] info = quantized_keys[key] assert spec is not None, f"Quantized key '{key}' in a checkpoint without a quantization spec" + global_scale = float(info.get("global_scale", 1.0)) + if not math.isfinite(global_scale) or global_scale <= 0: + raise ValueError(f"Invalid global scale for quantized H3 matrix {key}: {global_scale}") return QuantizedMatrix( weight=arrays[key], scales=arrays[f"{key}.scales"], biases=arrays[f"{key}.biases"] if info["has_biases"] else None, spec=spec, dequantized_dtype=_name_to_dtype(info["dequantized_dtype"]), + global_scale=global_scale, ) weights: dict[str, Any] = {} diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index 464250a9c2..ccc1bfe7cc 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -22,6 +22,7 @@ import json from dataclasses import dataclass from pathlib import Path +from typing import Any import numpy as np import mlx.core as mx @@ -77,9 +78,19 @@ def __init__(self, component_dir: Path): index_path = component_dir / "model.safetensors.index.json" self.key_to_shard: dict[str, str] = {} self._header_cache: dict[str, tuple[dict, int]] = {} + + def needed(key: str) -> bool: + if key == "model.language_model.embed_tokens.weight": + return True + prefix = "model.language_model.layers." + if not key.startswith(prefix): + return False + layer = key[len(prefix):].split(".", 1)[0] + return layer.isdigit() and int(layer) < TEXT_ENCODER_LAYER + if index_path.exists(): weight_map = json.loads(index_path.read_text())["weight_map"] - self.key_to_shard = {k: str(component_dir / s) for k, s in weight_map.items()} + self.key_to_shard = {k: str(component_dir / s) for k, s in weight_map.items() if needed(k)} else: single = component_dir / "model.safetensors" if not single.exists(): @@ -88,7 +99,7 @@ def __init__(self, component_dir: Path): with open(single, "rb") as handle: (header_len, ) = struct.unpack(" None: gc.collect() -_DTYPES = {"F32": np.float32, "F16": np.float16, "I64": np.int64, "I32": np.int32} +_DTYPES = {"F32": np.float32, "F16": np.float16, "I64": np.int64, "I32": np.int32, "U8": np.uint8} def _read_bf16_words(path: str, key: str, header: dict, data_start: int) -> np.ndarray: @@ -205,8 +216,20 @@ def _rms_norm(x, weight, eps: float): return x / mx.sqrt(mx.mean(x * x, axis=-1, keepdims=True) + eps) * weight +@dataclass(frozen=True) +class NVFP4Matrix: + """MLX row-major E2M1/E4M3 weights and the export's inverse global scale.""" + + weight: mx.array + scales: mx.array + global_scale: float + + def matmul(self, x): + return mx.quantized_matmul(x, self.weight, self.scales, mode="nvfp4") / self.global_scale + + def _linear(x, weight, bias=None): - y = x @ weight.T + y = weight.matmul(x) if isinstance(weight, NVFP4Matrix) else x @ weight.T if bias is not None: y = y + bias return y @@ -238,7 +261,7 @@ class StreamedMiniMaxH3TextConditioner: def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): self.component_dir = Path(component_dir) self.config = ConditionerConfig.from_config_json(self.component_dir / "config.json") - self.index = _ShardIndex(self.component_dir) + self.index: Any = _ShardIndex(self.component_dir) self.tokenizer = self._load_tokenizer(tokenizer_dir) def _load_tokenizer(self, tokenizer_dir: str | Path | None): @@ -284,28 +307,31 @@ def encode_tokens(self, token_ids: list[int]) -> tuple[np.ndarray, np.ndarray]: ]) cos, sin = _mrope_cos_sin(positions, cfg) - # Embedding rows gathered individually; the (151936, 5120) table is - # never fully materialized. - rows = [] - for token in token_ids: - key = "model.language_model.embed_tokens.weight" - rows.append(self.index.get_row(key, token)) - hidden = mx.array(np.stack(rows).astype(np.float32)) - del rows - gc.collect() + hidden = self._embed_tokens(token_ids) - if cfg.num_layers <= TEXT_ENCODER_LAYER: - raise ValueError(f"Conditioner needs > {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") + if cfg.num_layers < TEXT_ENCODER_LAYER: + raise ValueError(f"Conditioner needs at least {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") for layer in range(TEXT_ENCODER_LAYER): hidden = self._decoder_layer(layer, hidden, cos, sin) - # Per-layer sync: without this the whole 50-layer graph accumulates - # and the machine runs out of memory (same failure mode as the DiT). + # Per-layer sync keeps the 50-layer activation graph bounded. mx.eval(hidden) gc.collect() tags = np.full((seq_len, ), 1, dtype=np.int64) # MINIMAX_H3_TEXT_TAG return np.asarray(hidden).astype(np.float32), tags + def _embed_tokens(self, token_ids: list[int]): + # Embedding rows gathered individually; the full table is never + # materialized by the BF16 streaming path. + rows = [] + for token in token_ids: + key = "model.language_model.embed_tokens.weight" + rows.append(self.index.get_row(key, token)) + hidden = mx.array(np.stack(rows).astype(np.float32)) + del rows + gc.collect() + return hidden + # -- layers ---------------------------------------------------------- def _decoder_layer(self, index: int, hidden, cos, sin): @@ -368,6 +394,171 @@ def close(self) -> None: self.index.close() +def unswizzle_nvfp4_scales(scale: np.ndarray, rows: int, cols: int) -> np.ndarray: + """FlashInfer 128x4 scale bytes -> MLX row-major group-16 scale bytes.""" + pad_rows, pad_cols = -(-rows // 128) * 128, -(-cols // 4) * 4 + if scale.size != pad_rows * pad_cols: + raise ValueError(f"NVFP4 scales need {pad_rows * pad_cols} bytes, got {scale.size}.") + tiles = scale.reshape(pad_rows // 128, pad_cols // 4, 32, 4, 4) + return np.ascontiguousarray(tiles.transpose(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols]) + + +MLX_NVFP4_ENCODER_MANIFEST = "mlx_h3_nvfp4_encoder.json" + + +def _read_nvfp4_encoder_config(component_dir: str | Path) -> dict[str, Any]: + raw = json.loads((Path(component_dir) / "config.json").read_text()) + expected = { + "quant_method": "nvfp4", + "fmt": "e2m1", + "group_size": 16, + "scale_fmt": "e4m3", + "scale_layout": "128x4", + "activation_scheme": "dynamic", + } + quant = raw.get("quantization_config", {}) + if any(quant.get(key) != value for key, value in expected.items()): + raise ValueError("MLX NVFP4 conditioning requires the FastVideo group-16, 128x4 encoder export.") + return raw + + +def export_mlx_h3_nvfp4_encoder(component_dir: str | Path, output_dir: str | Path) -> Path: + """Cache the released packed encoder in MLX layout, without requantization. + + Keep original packed nibbles, row-major scale bytes, global scales and + embedding/norm values. Later loads skip CPU scale unswizzling and staging. + """ + raw = _read_nvfp4_encoder_config(component_dir) + output_dir = Path(output_dir) + if output_dir.exists() and any(output_dir.iterdir()): + raise FileExistsError(f"Encoder cache output must be empty: {output_dir}") + output_dir.mkdir(parents=True, exist_ok=True) + index = _ResidentNVFP4Index(_ShardIndex(Path(component_dir))) + try: + arrays = {} + matrices = {} + dense_keys = [] + for key, value in index.weights.items(): + if isinstance(value, NVFP4Matrix): + arrays[key] = value.weight + arrays[key + ".scales"] = value.scales + matrices[key] = {"global_scale": value.global_scale} + else: + arrays[key] = value + dense_keys.append(key) + mx.save_safetensors(str(output_dir / "model.safetensors"), arrays) + (output_dir / "config.json").write_text(json.dumps(raw, indent=2) + "\n") + manifest = { + "format_version": 1, + "language_layers": TEXT_ENCODER_LAYER, + "matrices": matrices, + "dense_keys": dense_keys, + "source_dir": str(Path(component_dir).resolve()) + } + (output_dir / MLX_NVFP4_ENCODER_MANIFEST).write_text(json.dumps(manifest, indent=2) + "\n") + finally: + index.close() + return output_dir + + +class _ResidentNVFP4Index: + """Load the released 50-layer encoder without expanding packed matrices.""" + + def __init__(self, source: _ShardIndex): + self.weights: dict[str, mx.array | NVFP4Matrix] = {} + for key in sorted(source.key_to_shard): + if key.endswith(".weight_packed"): + prefix = key.removesuffix(".weight_packed") + packed = np.array(source.get(key), copy=True) + if packed.dtype != np.uint8 or packed.ndim != 2 or packed.shape[1] % 4: + raise ValueError(f"Invalid packed NVFP4 matrix {key}: {packed.shape}, {packed.dtype}") + rows, cols = packed.shape[0], packed.shape[1] * 2 + if cols % 16: + raise ValueError(f"NVFP4 input width must be divisible by 16: {key}") + scales = unswizzle_nvfp4_scales(source.get(prefix + ".weight_scale"), rows, cols // 16) + global_scale = float(source.get(prefix + ".weight_global_scale").reshape(-1)[0]) + if not np.isfinite(global_scale) or global_scale <= 0: + raise ValueError(f"Invalid NVFP4 global scale for {prefix}: {global_scale}") + weight = mx.array(packed).view(mx.uint32) + scale_bytes = mx.array(scales) + if not bool(mx.all(mx.isfinite(mx.from_fp8(scale_bytes, dtype=mx.float32)))): + raise ValueError(f"Non-finite NVFP4 block scales for {prefix}") + mx.eval(weight, scale_bytes) + self.weights[prefix + ".weight"] = NVFP4Matrix(weight, scale_bytes, global_scale) + elif key.endswith(".weight"): + if key == "model.language_model.embed_tokens.weight": + shard = source.key_to_shard[key] + header, data_start = source._cache_header(shard) + if header[key]["dtype"] == "BF16": + value = mx.array(_read_bf16_words(shard, key, header, data_start)).view(mx.bfloat16) + else: + value = mx.array(source.get(key)) + else: + value = source.get_mlx(key) + mx.eval(value) + self.weights[key] = value + source.close() + + @classmethod + def from_mlx_checkpoint(cls, component_dir: str | Path): + component_dir = Path(component_dir) + manifest = json.loads((component_dir / MLX_NVFP4_ENCODER_MANIFEST).read_text()) + if manifest.get("format_version") != 1 or manifest.get("language_layers") != TEXT_ENCODER_LAYER: + raise ValueError("Unsupported native MLX NVFP4 encoder cache") + arrays = mx.load(str(component_dir / "model.safetensors")) + matrices = manifest["matrices"] + expected = set(manifest["dense_keys"]) | set(matrices) | {key + ".scales" for key in matrices} + if set(arrays) != expected: + raise ValueError("Native MLX encoder arrays do not match the manifest") + index = cls.__new__(cls) + index.weights = {key: arrays[key] for key in manifest["dense_keys"]} + for key, info in matrices.items(): + weight, scales = arrays[key], arrays[key + ".scales"] + factor = float(info["global_scale"]) + if (weight.ndim != 2 or weight.dtype != mx.uint32 or scales.dtype != mx.uint8 + or scales.shape != (weight.shape[0], weight.shape[1] // 2) or weight.shape[1] % 2 + or not np.isfinite(factor) or factor <= 0): + raise ValueError(f"Invalid native MLX NVFP4 matrix: {key}") + index.weights[key] = NVFP4Matrix(weight, scales, factor) + mx.eval(list(arrays.values())) + return index + + def get_mlx(self, key: str): + return self.weights[key] + + def close(self) -> None: + self.weights.clear() + gc.collect() + + +class ResidentNVFP4MiniMaxH3TextConditioner(StreamedMiniMaxH3TextConditioner): + """Released NVFP4 encoder weights with floating-point MLX activations. + + The packed weights and embedding table stay resident. CUDA quantizes + activations to FP4; this path keeps FP32 activations, so hidden states are + not expected to be bit-exact with the CUDA encoder. + """ + + def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): + _read_nvfp4_encoder_config(component_dir) + # Fail on an older MLX before reading the encoder's large shards. + try: + packed, scales = mx.quantize(mx.ones((1, 64)), mode="nvfp4") + mx.eval(mx.quantized_matmul(mx.ones((1, 64)), packed, scales, mode="nvfp4")) + except (ValueError, RuntimeError) as error: + raise RuntimeError("Native NVFP4 conditioning requires an MLX build with nvfp4 matmul support.") from error + super().__init__(component_dir, tokenizer_dir) + if (Path(component_dir) / MLX_NVFP4_ENCODER_MANIFEST).exists(): + self.index.close() + self.index = _ResidentNVFP4Index.from_mlx_checkpoint(component_dir) + else: + self.index = _ResidentNVFP4Index(self.index) + + def _embed_tokens(self, token_ids: list[int]): + table = self.index.get_mlx("model.language_model.embed_tokens.weight") + return table[mx.array(token_ids, dtype=mx.int32)].astype(mx.float32) + + def _apply_mrope(q_or_k, cos, sin): """q_or_k: (S, H, D); cos/sin: (S, 1, D).""" half = q_or_k.shape[-1] // 2 diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index c6ca112f32..be898f3bb8 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -49,6 +49,7 @@ MINIMAX_H3_MIN_DURATION, MINIMAX_H3_VIDEO_SHIFT, MiniMaxH3SchedulerState, + _eval_value, adaln_timestep_union, align_num_frames, audio_latent_num_frames, @@ -231,7 +232,7 @@ def _cleanup_mlx() -> None: def _default_metal_wired_limit_gib(mx) -> float: - """Keep the default below both physical memory and the tested 30 GiB cap.""" + """Legacy helper for allocator capacity, not the wired-residency setting.""" metal = getattr(mx, "metal", None) if metal is None: return 30.0 @@ -244,6 +245,30 @@ def _default_metal_wired_limit_gib(mx) -> float: return min(30.0, 0.84 * total_bytes / 2**30) +def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> None: + """Keep allocator capacity separate from explicitly requested wired residency.""" + set_memory = getattr(mx, "set_memory_limit", None) + if set_memory is None and hasattr(mx, "metal"): + set_memory = getattr(mx.metal, "set_memory_limit", None) + if set_memory is not None: + try: + set_memory(int(_default_metal_wired_limit_gib(mx) * 2**30)) + except Exception as error: # noqa: BLE001 - older MLX best effort + logger.info("Could not set the Metal allocation limit: %s", error) + if wired_limit_gib is None: + return + if not math.isfinite(wired_limit_gib) or wired_limit_gib <= 0: + raise ValueError("metal_wired_limit_gib must be finite and positive") + set_wired = getattr(mx, "set_wired_limit", None) + if set_wired is None and hasattr(mx, "metal"): + set_wired = getattr(mx.metal, "set_wired_limit", None) + if set_wired is None: + raise RuntimeError("This MLX build cannot set the requested wired-memory limit") + # Explicit requests must succeed; do not silently benchmark an unwired model. + previous = set_wired(int(wired_limit_gib * 2**30)) + logger.info("MLX wired limit %.2f GiB (previous %.2f GiB)", wired_limit_gib, previous / 2**30) + + MINIMAX_H3_PROMPT_CACHE_VERSION = "v2-attention-layout" @@ -328,21 +353,20 @@ def __init__( video_decode_backend: str = "h3-vae", taeh3_checkpoint: str | Path | None = None, taeh3_chunk_size: int = 5, + conditioner_mode: str = "auto", + resident: bool = False, ) -> None: import mlx.core as mx - set_limit = getattr(mx, "set_memory_limit", None) - if set_limit is None and hasattr(mx, "metal"): - set_limit = getattr(mx.metal, "set_memory_limit", None) - if set_limit is not None: - # Keep large resident models inside a predictable wired budget. - try: - if metal_wired_limit_gib is None: - metal_wired_limit_gib = _default_metal_wired_limit_gib(mx) - set_limit(int(metal_wired_limit_gib * 2**30)) - except Exception as error: # noqa: BLE001 - best effort on older MLX - logger.info("Could not raise the Metal wired limit: %s", error) + _configure_metal_memory_limits(mx, metal_wired_limit_gib) self.model_root = Path(model_root) + if conditioner_mode not in ("auto", "streamed", "nvfp4"): + raise ValueError(f"Unknown H3 conditioner mode: {conditioner_mode}") + if resident and video_decode_backend != "h3-vae": + raise ValueError("Resident H3 generation requires the H3 video VAE.") + self.conditioner_mode = conditioner_mode + self.resident = resident + self._resident_components: dict[str, Any] = {} self.dit_checkpoint = Path(mlx_dit_checkpoint) self.vae_dtype = vae_dtype if video_decode_backend not in ("h3-vae", "taeh3"): @@ -414,20 +438,70 @@ def resolve_geometry( # -- phase 1: conditioning ------------------------------------------- + def prepare_resident(self) -> None: + """Load the encoder, DiT, and both decoders once, before timed requests.""" + if not self.resident or self._resident_components: + return + from fastvideo.mlx_runtime.minimax_h3_conditioner import ResidentNVFP4MiniMaxH3TextConditioner + from fastvideo.mlx_runtime.minimax_h3_audio_vae import mlx_h3_audio_vae_from_dir + from fastvideo.mlx_runtime.minimax_h3_video_vae import mlx_h3_video_vae_from_dir + + import mlx.core as mx + + try: + conditioner = self._load_conditioner() + if not isinstance(conditioner, ResidentNVFP4MiniMaxH3TextConditioner): + conditioner.close() + raise ValueError("All-resident generation requires the packed NVFP4 text encoder.") + self._resident_components["conditioner"] = conditioner + dit = load_mlx_h3_checkpoint(self.dit_checkpoint) + self._resident_components["dit"] = dit + for group in [dit.weights, *dit.blocks, *dit.refiner]: + for value in group.values(): + _eval_value(value) + cache = dit._adaln_cache + if cache is not None: + mx.eval(cache.block_tables, cache.norm_out_shift, cache.norm_out_scale) + self._resident_components["video_vae"] = mlx_h3_video_vae_from_dir(self.model_root / "vae", + include_encoder=False, + storage_dtype=self.vae_dtype) + self._resident_components["audio_vae"] = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", + include_encoder=False) + mx.eval(list(self._resident_components["audio_vae"].weights.values())) + logger.info("H3 components resident: %.2f GiB active MLX memory", mx.get_active_memory() / 2**30) + except Exception: + self.close() + raise + + def close(self) -> None: + conditioner = self._resident_components.get("conditioner") + if conditioner is not None: + conditioner.close() + self._resident_components.clear() + _cleanup_mlx() + def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: """Returns (hidden states (S, hidden), token tags). Uses the cache or the streamed conditioner.""" cache_key = None if self.prompt_cache_dir is not None: - cache_key = prompt_cache_path(self.prompt_cache_dir, self.model_root, prompt) + identity = (f"{self.model_root}::conditioner=" + f"{getattr(self, 'conditioner_dir', self.model_root / 'text_encoder')}::" + f"{getattr(self, 'conditioner_mode', 'auto')}") + cache_key = prompt_cache_path(self.prompt_cache_dir, identity, prompt) if cache_key.exists(): data = np.load(cache_key) logger.info("Loaded prompt embeddings from cache %s", cache_key) return data["hidden_states"], data["token_tags"] - conditioner = self._load_conditioner() + if getattr(self, "resident", False): + self.prepare_resident() + conditioner = self._resident_components["conditioner"] + else: + conditioner = self._load_conditioner() hidden, tags = conditioner.encode_prompt(prompt) - conditioner.close() + if not getattr(self, "resident", False): + conditioner.close() _cleanup_mlx() if cache_key is not None: cache_key.parent.mkdir(parents=True, exist_ok=True) @@ -450,8 +524,17 @@ def has_conditioner_weights(self) -> bool: return marker.exists() or single.exists() def _load_conditioner(self): - from fastvideo.mlx_runtime.minimax_h3_conditioner import StreamedMiniMaxH3TextConditioner + from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + ResidentNVFP4MiniMaxH3TextConditioner, + StreamedMiniMaxH3TextConditioner, + ) + config = json.loads((self.conditioner_dir / "config.json").read_text()) + packed = config.get("quantization_config", {}).get("quant_method") == "nvfp4" + if self.conditioner_mode == "nvfp4" or (self.conditioner_mode == "auto" and packed): + return ResidentNVFP4MiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + if packed: + raise ValueError("The streamed conditioner requires BF16 weights; use conditioner_mode='nvfp4'.") return StreamedMiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) # -- phase 2: denoise -------------------------------------------------- @@ -478,6 +561,10 @@ def denoise( geometry = self.resolve_geometry(height, width, num_frames, enforce_duration=audio_num_frames is None) audio_frames = geometry["num_frames"] if audio_num_frames is None else align_num_frames(audio_num_frames) + if dit is None and getattr(self, "resident", False): + _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) + self.prepare_resident() + dit = self._resident_components["dit"] owned_dit = dit is None if owned_dit: _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) @@ -624,7 +711,13 @@ def decode_video(self, raise RuntimeError(f"TAEH3 produced unexpected frame shape: {frames.shape}") _cleanup_mlx() return frames - vae = mlx_h3_video_vae_from_dir(self.model_root / "vae", include_encoder=False, storage_dtype=self.vae_dtype) + if getattr(self, "resident", False): + self.prepare_resident() + vae = self._resident_components["video_vae"] + else: + vae = mlx_h3_video_vae_from_dir(self.model_root / "vae", + include_encoder=False, + storage_dtype=self.vae_dtype) expected_height = height // vae.spatial_compression_ratio expected_width = width // vae.spatial_compression_ratio if (geometry["latent_height"], geometry["latent_width"]) != (expected_height, expected_width): @@ -666,7 +759,11 @@ def decode_audio(self, audio_rows: np.ndarray, *, num_frames: int) -> np.ndarray num_audio_latents = audio_latent_num_frames(align_num_frames(num_frames)) latents = unpack_audio_tokens(audio_rows, num_audio_latents) - vae = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", include_encoder=False) + if getattr(self, "resident", False): + self.prepare_resident() + vae = self._resident_components["audio_vae"] + else: + vae = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", include_encoder=False) z = vae.denormalize_latents(mx.array(latents)) waveform = np.asarray(vae.decode(z))[:, 0, :] # (B, 1, S) -> (B, S) del vae, z diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa.py b/fastvideo/mlx_runtime/minimax_h3_vsa.py index 3c49969e5c..48ca15fa32 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa.py @@ -544,11 +544,11 @@ def _key_valid_mask(block_idx, variable_block_sizes, tile_elems: int): return offsets[None, None, None, :] < selected_sizes[:, :, :, None] -_REFERENCE_GATHER_TARGET_BYTES = 2 * 1024**3 +_REFERENCE_GATHER_TARGET_BYTES = 256 * 1024**2 def _reference_gather_query_chunk(heads: int, dim: int, k_sel: int, tile_elems: int, n_q: int) -> int: - """Batch as many query tiles as fit in ~2 GiB of gathered BF16 K/V.""" + """Bound gathered BF16 K/V to 256 MiB, leaving space for resident weights.""" bytes_per_query = 4 * heads * max(k_sel, 1) * tile_elems * dim chunk = min(n_q, max(1, _REFERENCE_GATHER_TARGET_BYTES // max(bytes_per_query, 1))) return int(chunk) diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index bb21b1912a..14068f14de 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -37,16 +37,18 @@ """ # One threadgroup = one (head, video query tile). 8 SIMD-groups x 32 = 256 -# threads cover 64 query rows. Q is half smem (16 KiB). K/V stage 8 keys. +# threads cover 64 query rows. K/V stage 32 keys, with 28.25 KiB total smem. +# Updating online softmax once per 32 keys reduces accumulator rescaling barriers. +# Four SIMD lanes cooperate per query row in the softmax reduction. _SIMD_SOURCE = """ const int TILE = 64; const int D = 128; const int SG = 32; const int N_SG = 8; const int ROWS = 8; - const int KCHUNK = 8; - threadgroup float kvsmem[8 * 128]; - threadgroup float score_smem[8 * 64]; + const int KCHUNK = 32; + threadgroup float kvsmem[KCHUNK * D]; + threadgroup float score_smem[N_SG * ROWS * KCHUNK]; threadgroup float scale_tmp[8 * 64]; threadgroup float qtile[8 * 64]; threadgroup float row_alpha[8 * 8]; @@ -65,7 +67,7 @@ int qt = n_prefix + (int)q_tile; int q_valid = active ? vbs[qt] : 0; int q_base_tile = (((int)head * S) + qt * TILE) * D; - threadgroup float *sg_scores = score_smem + sid * 64; + threadgroup float *sg_scores = score_smem + sid * ROWS * KCHUNK; threadgroup float *sg_tmp = scale_tmp + sid * 64; threadgroup float *sg_qtile = qtile + sid * 64; threadgroup float *sg_alpha = row_alpha + sid * 8; @@ -117,34 +119,38 @@ } threadgroup_barrier(mem_flags::mem_threadgroup); - thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 kmat; - simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kk * 8), D, ulong2(0, 0), true); - simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); + for (int kc = 0; kc < KCHUNK; kc += 8) { + thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 kmat; + simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D, ulong2(0, 0), true); + simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); + } + simdgroup_store(smat, sg_scores + kc, KCHUNK); } - simdgroup_store(smat, sg_scores, KCHUNK); simdgroup_barrier(mem_flags::mem_threadgroup); - float scores[8]; + float scores[KCHUNK / 4]; float cmax = -3.402823466e+38f; - if (lane < (uint)ROWS) { - int grow = qrow0 + (int)lane; - for (int t = 0; t < KCHUNK; t++) { - int gtok = j0 + t; - float sc = -3.402823466e+38f; - if (grow < q_valid && gtok < k_valid && gtok < TILE) { - sc = sg_scores[(int)lane * KCHUNK + t] * scale; - } - scores[t] = sc; - cmax = metal::max(cmax, sc); + int row = (int)lane / 4; + int col_lane = (int)lane % 4; + int grow = qrow0 + row; + for (int t = col_lane; t < KCHUNK; t += 4) { + int gtok = j0 + t; + float sc = -3.402823466e+38f; + if (grow < q_valid && gtok < k_valid && gtok < TILE) { + sc = sg_scores[row * KCHUNK + t] * scale; } - float m_new = metal::max(row_m, cmax); - float alpha = metal::exp(row_m - m_new); - row_lse *= alpha; - row_m = m_new; - sg_alpha[(int)lane] = alpha; + scores[t / 4] = sc; + cmax = metal::max(cmax, sc); } + cmax = metal::max(cmax, simd_shuffle_xor(cmax, 1)); + cmax = metal::max(cmax, simd_shuffle_xor(cmax, 2)); + float m_new = metal::max(row_m, cmax); + float alpha = metal::exp(row_m - m_new); + row_lse *= alpha; + row_m = m_new; + if (col_lane == 0) sg_alpha[row] = alpha; simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { scale_rows_simd8x8(acc[kk], sg_alpha, sg_tmp, lane); @@ -164,33 +170,32 @@ threadgroup_barrier(mem_flags::mem_threadgroup); float local = 0.0f; - if (lane < (uint)ROWS) { - int grow = qrow0 + (int)lane; - for (int t = 0; t < KCHUNK; t++) { - float w = 0.0f; - if (grow < q_valid) { - w = metal::exp(scores[t] - row_m); - } - sg_scores[(int)lane * KCHUNK + t] = w; - local += w; - } - row_lse += local; + for (int t = col_lane; t < KCHUNK; t += 4) { + float w = 0.0f; + if (grow < q_valid) w = metal::exp(scores[t / 4] - row_m); + sg_scores[row * KCHUNK + t] = w; + local += w; } + local += simd_shuffle_xor(local, 1); + local += simd_shuffle_xor(local, 2); + row_lse += local; simdgroup_barrier(mem_flags::mem_threadgroup); - thread simdgroup_float8x8 pmat; - simdgroup_load(pmat, sg_scores, KCHUNK); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 vmat; - simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kk * 8), D); - simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); + for (int kc = 0; kc < KCHUNK; kc += 8) { + thread simdgroup_float8x8 pmat; + simdgroup_load(pmat, sg_scores + kc, KCHUNK); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 vmat; + simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D); + simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); + } } threadgroup_barrier(mem_flags::mem_threadgroup); } } - if (lane < (uint)ROWS) { - sg_alpha[(int)lane] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; + if (lane % 4 == 0) { + sg_alpha[(int)lane / 4] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; } simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { @@ -211,6 +216,7 @@ } simdgroup_barrier(mem_flags::mem_threadgroup); } + """ diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py new file mode 100644 index 0000000000..caac1541e8 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exercise native floating-point quantized H3 storage on supported MLX builds.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +from fastvideo.mlx_runtime.fastwan import MLXQuantizationSpec, ensure_quantization_supported, linear +from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, quantize_matrix, save_mlx_h3_checkpoint + + +@pytest.mark.parametrize('mode', ['mxfp8', 'mxfp4', 'nvfp4']) +def test_float_quantized_checkpoint_preserves_matrix(tmp_path, mode): + spec = MLXQuantizationSpec.from_name(mode) + ensure_quantization_supported(spec) + dense = (mx.random.normal((64, 64)) * 0.001).astype(mx.bfloat16) + weight = quantize_matrix(dense, spec) + restored = mx.dequantize(weight.weight, weight.scales, mode=mode).astype(mx.float32) * weight.global_scale + relative_error = mx.sqrt(mx.sum((restored - dense.astype(mx.float32))**2) / mx.sum(dense.astype(mx.float32)**2)) + assert float(relative_error.item()) < 0.15 + x = mx.random.normal((3, 64)).astype(mx.bfloat16) + config = dict(hidden_size=64, num_attention_heads=1, attention_head_dim=64, ffn_dim=128, + in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=64, + freq_dim=64, time_embed_dim=64, rope_freq_dim=4, rope_theta=10000., + norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5) + dit = MLXMiniMaxH3DiT({'test.weight': weight}, [], [], config) + save_mlx_h3_checkpoint(dit, tmp_path) + loaded = load_mlx_h3_checkpoint(tmp_path) + actual = linear(x, loaded.weights['test.weight']).astype(mx.float32) + expected = linear(x, weight).astype(mx.float32) + np.testing.assert_array_equal(np.array(actual), np.array(expected)) + assert loaded.weights['test.weight'].biases is None + assert loaded.weights['test.weight'].spec == spec + assert loaded.weights['test.weight'].global_scale == weight.global_scale diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py new file mode 100644 index 0000000000..cd954981b6 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py @@ -0,0 +1,72 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Check the pruned model's shared rank-16 modulation against NumPy math.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +from fastvideo.mlx_runtime.fastwan import timestep_embedding +from fastvideo.mlx_runtime.minimax_h3 import ( + MLXMiniMaxH3DiT, + MiniMaxH3SchedulerState, + adaln_timestep_union, + load_mlx_h3_checkpoint, + save_mlx_h3_checkpoint, +) + + +def test_rank16_cache_has_one_silu_before_shared_basis(tmp_path): + rng = np.random.default_rng(2026) + hidden, rank = 32, 16 + + def array(shape): + return mx.array(rng.normal(0, 0.1, shape).astype(np.float32)) + + config = dict(hidden_size=hidden, num_attention_heads=1, attention_head_dim=hidden, ffn_dim=64, + in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=hidden, + freq_dim=hidden, time_embed_dim=hidden, rope_freq_dim=4, rope_theta=10000., + norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5, adaln_rank=rank, num_layers=42) + weights = { + 'time_embedder.linear_1.weight': array((hidden, hidden)), + 'time_embedder.linear_1.bias': array((hidden,)), + 'time_embedder.linear_2.weight': array((hidden, hidden)), + 'time_embedder.linear_2.bias': array((hidden,)), + 'adaln_basis.weight': array((rank, hidden)), + 'norm_out.linear.weight': array((2 * hidden, rank)), + 'norm_out.linear.bias': array((2 * hidden,)), + } + blocks = [{'attn.to_q.weight': array((hidden, hidden)), + 'adaln_proj.linear.weight': array((18 * hidden, rank)), + 'adaln_proj.linear.bias': array((18 * hidden,))} for _ in range(42)] + dit = MLXMiniMaxH3DiT(weights, blocks, [], config) + rungs = [999, 874, 749, 624, 500, 375, 250, 125] + timesteps = adaln_timestep_union(MiniMaxH3SchedulerState.from_dmd_steps(10, rungs), + MiniMaxH3SchedulerState.from_dmd_steps(3, rungs)) + + def project(x, weight, bias=None): + result = x @ np.array(weight).T + return result if bias is None else result + np.array(bias) + + def silu(x): + return x / (1 + np.exp(-x)) + + features = np.array(timestep_embedding(mx.array(timesteps), hidden)) + first = project(features, weights['time_embedder.linear_1.weight'], weights['time_embedder.linear_1.bias']) + second = project(silu(first), weights['time_embedder.linear_2.weight'], weights['time_embedder.linear_2.bias']) + expected_basis = project(silu(second), weights['adaln_basis.weight']) + np.testing.assert_allclose(np.array(dit.compute_temb(mx.array(timesteps))), expected_basis, atol=1e-6) + expected_blocks = [project(expected_basis, block['adaln_proj.linear.weight'], + block['adaln_proj.linear.bias']).reshape(-1, 6 * hidden) for block in blocks] + expected_out = project(expected_basis, weights['norm_out.linear.weight'], weights['norm_out.linear.bias']) + cache = dit.precompute_adaln(timesteps) + for tables, expected in zip(cache.block_tables, expected_blocks, strict=True): + np.testing.assert_allclose(np.concatenate([np.array(t) for t in tables], axis=-1), expected, atol=1e-6) + np.testing.assert_allclose(np.array(cache.norm_out_shift), expected_out[:, :hidden], atol=1e-6) + np.testing.assert_allclose(np.array(cache.norm_out_scale), expected_out[:, hidden:], atol=1e-6) + assert all(block['adaln_proj.linear.weight'] is None for block in blocks) + + save_mlx_h3_checkpoint(dit, tmp_path) + loaded = load_mlx_h3_checkpoint(tmp_path) + assert loaded.adaln_rank == rank + assert len(loaded.blocks) == 42 + np.testing.assert_array_equal(loaded._adaln_cache.timesteps, timesteps) + np.testing.assert_array_equal(np.array(loaded._adaln_cache.norm_out_scale), np.array(cache.norm_out_scale)) diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py new file mode 100644 index 0000000000..288b3bf465 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py @@ -0,0 +1,176 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Native MLX NVFP4 encoder storage and residency regression checks.""" +import json +from types import SimpleNamespace + +import numpy as np +import pytest + +mx = pytest.importorskip("mlx.core") +from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + NVFP4Matrix, ResidentNVFP4MiniMaxH3TextConditioner, _ResidentNVFP4Index, + _ShardIndex, export_mlx_h3_nvfp4_encoder, unswizzle_nvfp4_scales, +) +from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline + + +def _swizzle(values): + # Independent coordinate mapping of FlashInfer's 128-row, four-column tiles. + rows, cols = values.shape + padded = np.zeros((-(-rows // 128) * 128, -(-cols // 4) * 4), np.uint8) + padded[:rows, :cols] = values + output = np.empty(padded.size, np.uint8) + for r in range(padded.shape[0]): + for c in range(padded.shape[1]): + address = ((((r // 128) * (padded.shape[1] // 4) + c // 4) * 32 + + r % 32) * 4 + (r % 128) // 32) * 4 + c % 4 + output[address] = padded[r, c] + return output + + +def _decode_e4m3(values): + sign = np.where(values & 128, -1.0, 1.0) + exponent = (values >> 3) & 15 + fraction = values & 7 + return sign * np.where(exponent == 0, fraction * 2.0**-9, + (1.0 + fraction / 8.0) * 2.0**(exponent.astype(int) - 7)) + + +def test_padded_scale_layout_round_trip(): + rng = np.random.default_rng(11) + values = rng.integers(0, 127, (140, 7), dtype=np.uint8) + np.testing.assert_array_equal(unswizzle_nvfp4_scales(_swizzle(values), 140, 7), values) + with pytest.raises(ValueError, match="bytes"): + unswizzle_nvfp4_scales(np.zeros(1, np.uint8), 140, 7) + + +@pytest.mark.parametrize("global_scale", [0.5, 4.0]) +@pytest.mark.parametrize("native_cache", [False, True]) +def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, global_scale, native_cache): + from safetensors.numpy import save_file + + rng = np.random.default_rng(3) + packed = rng.integers(0, 256, (128, 64), dtype=np.uint8) + scales = rng.integers(24, 96, (128, 8), dtype=np.uint8) + prefix = "model.language_model.layers.0.self_attn.q_proj" + save_file({prefix + ".weight_packed": packed, + prefix + ".weight_scale": _swizzle(scales), + prefix + ".weight_global_scale": np.array([global_scale], np.float32)}, + tmp_path / "model.safetensors") + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + if native_cache: + _write_encoder_config(tmp_path) + cache_dir = export_mlx_h3_nvfp4_encoder(tmp_path, tmp_path / "cache") + cached = _ResidentNVFP4Index.from_mlx_checkpoint(cache_dir) + for key, original in index.weights.items(): + value = cached.get_mlx(key) + np.testing.assert_array_equal(np.array(value.weight), np.array(original.weight)) + np.testing.assert_array_equal(np.array(value.scales), np.array(original.scales)) + assert value.global_scale == original.global_scale + index.close() + index = cached + weight = index.get_mlx(prefix + ".weight") + assert isinstance(weight, NVFP4Matrix) + assert weight.weight.dtype == mx.uint32 + assert weight.scales.dtype == mx.uint8 + lut = np.array([0, .5, 1, 1.5, 2, 3, 4, 6, 0, -.5, -1, -1.5, -2, -3, -4, -6], np.float32) + dense = np.stack((lut[packed & 15], lut[packed >> 4]), axis=-1).reshape(128, 128) + dense *= np.repeat(_decode_e4m3(scales), 16, axis=1) / global_scale + x = rng.standard_normal((3, 128)).astype(np.float32) + np.testing.assert_allclose(np.array(weight.matmul(mx.array(x))), x @ dense.T, rtol=3e-5, atol=1e-3) + index.close() + assert not index.weights + + +def _write_encoder_config(path): + (path / "config.json").write_text(json.dumps({"quantization_config": { + "quant_method": "nvfp4", "fmt": "e2m1", "group_size": 16, + "scale_fmt": "e4m3", "scale_layout": "128x4", "activation_scheme": "dynamic", + }})) + + +@pytest.mark.parametrize("native_cache", [False, True]) +def test_resident_embedding_keeps_bf16_storage(tmp_path, native_cache): + torch = pytest.importorskip("torch") + from safetensors.torch import save_file + + key = "model.language_model.embed_tokens.weight" + table = torch.arange(60).reshape(10, 6).to(torch.bfloat16) + save_file({key: table}, tmp_path / "model.safetensors") + if native_cache: + _write_encoder_config(tmp_path) + cache_dir = export_mlx_h3_nvfp4_encoder(tmp_path, tmp_path / "cache") + index = _ResidentNVFP4Index.from_mlx_checkpoint(cache_dir) + with pytest.raises(FileExistsError, match="empty"): + export_mlx_h3_nvfp4_encoder(tmp_path, cache_dir) + else: + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + assert index.get_mlx(key).dtype == mx.bfloat16 + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + conditioner.index = index + np.testing.assert_array_equal(np.array(conditioner._embed_tokens([7, 1])), table[[7, 1]].float().numpy()) + conditioner.close() + + +def _pipeline(): + pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) + pipeline.resident = True + pipeline._resident_components = {} + pipeline.dit_checkpoint = "tiny" + pipeline.model_root = __import__("pathlib").Path("tiny") + pipeline.vae_dtype = "fp16" + return pipeline + + +def test_resident_preload_reuses_models_and_evaluates_audio(monkeypatch): + import fastvideo.mlx_runtime.minimax_h3_pipeline as module + import fastvideo.mlx_runtime.minimax_h3_audio_vae as audio + import fastvideo.mlx_runtime.minimax_h3_video_vae as video + + pipeline = _pipeline() + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + conditioner.index = SimpleNamespace(close=lambda: None) + monkeypatch.setattr(pipeline, "_load_conditioner", lambda: conditioner) + calls = [] + def load_dit(path): + calls.append(path) + return SimpleNamespace(weights={"x": mx.ones((1,))}, blocks=[], refiner=[], _adaln_cache=None) + monkeypatch.setattr(module, "load_mlx_h3_checkpoint", load_dit) + monkeypatch.setattr(video, "mlx_h3_video_vae_from_dir", lambda *a, **k: object()) + decoder = SimpleNamespace(weights={"x": mx.ones((2,))}) + monkeypatch.setattr(audio, "mlx_h3_audio_vae_from_dir", lambda *a, **k: decoder) + pipeline.prepare_resident() + pipeline.prepare_resident() + assert calls == ["tiny"] + assert set(pipeline._resident_components) == {"conditioner", "dit", "video_vae", "audio_vae"} + pipeline.close() + assert not pipeline._resident_components + + +def test_failed_preload_releases_encoder(monkeypatch): + import fastvideo.mlx_runtime.minimax_h3_pipeline as module + + pipeline = _pipeline() + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + closed = [] + conditioner.index = SimpleNamespace(close=lambda: closed.append(True)) + monkeypatch.setattr(pipeline, "_load_conditioner", lambda: conditioner) + def fail(path): + raise RuntimeError("out of memory") + monkeypatch.setattr(module, "load_mlx_h3_checkpoint", fail) + with pytest.raises(RuntimeError, match="out of memory"): + pipeline.prepare_resident() + assert closed == [True] + assert not pipeline._resident_components + + +def test_single_shard_omits_unused_language_layers_and_vision(tmp_path): + from safetensors.numpy import save_file + + kept = "model.language_model.layers.49.input_layernorm.weight" + dropped = "model.language_model.layers.50.input_layernorm.weight" + save_file({kept: np.ones(8, np.float32), dropped: np.ones(8, np.float32), + "model.visual.weight": np.ones((8, 8), np.float32)}, tmp_path / "model.safetensors") + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + assert set(index.weights) == {kept} + index.close() diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py new file mode 100644 index 0000000000..376478624a --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py @@ -0,0 +1,24 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Changing the gather memory budget must preserve selected-tile attention.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +import fastvideo.mlx_runtime.minimax_h3_vsa as vsa + + +def test_gather_budget_preserves_attention_with_partial_tiles(monkeypatch): + geom = vsa.build_h3_tile_geometry((7, 67), (6, 8, 8), 64) + rng = np.random.default_rng(2026) + shape = (geom.padded_length, 2, 128) + q, k, value = [mx.array(rng.normal(size=shape).astype(np.float32)).astype(mx.bfloat16) for _ in range(3)] + # Reverse video order also exercises the selected-key order, not a dense mask. + selected = np.array([0, 1, geom.num_tiles - 1, geom.num_prefix_tiles], dtype=np.int32) + idx = mx.array(np.broadcast_to(selected, (2, geom.num_video_tiles, selected.size)).copy()) + monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 2 * 1024**3) + expected = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) + mx.eval(expected) + monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 1) + actual = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) + mx.eval(actual) + np.testing.assert_array_equal(np.array(actual.astype(mx.float32)), np.array(expected.astype(mx.float32))) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index b63b50a067..e4d08f830f 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -5,6 +5,8 @@ # Copyright 2025 The FastVideo Authors. from __future__ import annotations + +import os import contextlib import re from collections.abc import Callable, Generator diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index 21be0864a9..e4f7a08bf3 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -16,6 +16,7 @@ import torch.nn.functional as F from torch.utils.checkpoint import checkpoint +import fastvideo.envs as envs from fastvideo.attention import get_attn_backend from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig from fastvideo.platforms import AttentionBackendEnum @@ -525,7 +526,7 @@ 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"))) + return max(1, envs.FASTVIDEO_H3_VAE_TILE_BATCH.get()) def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool: diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 7dbb74d1c1..0bd2641e5f 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -466,6 +466,9 @@ def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: def _move_module(self, module: Any, device: str | torch.device) -> bool: if _module_has_dtensor_params(module): return False + if not callable(getattr(module, "named_parameters", None)): + module.to(device) + return True if os.environ.get("FASTVIDEO_H3_PINNED_SWAP", "1") == "1": _pinned_swap(module, torch.device(device)) else: diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index c87d946f65..71d9dd0133 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -191,8 +191,14 @@ def _execute( 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) + stage_name, + torch.cuda.max_memory_allocated() / gib, + torch.cuda.memory_reserved() / gib, + torch.cuda.memory_allocated() / gib) + batch.logging_info.add_stage_metric(stage_key, "peak_allocated_mb", + torch.cuda.max_memory_allocated() / 1024**2) + batch.logging_info.add_stage_metric(stage_key, "peak_reserved_mb", + torch.cuda.max_memory_reserved() / 1024**2) 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) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py index e009505547..6e59682545 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py @@ -21,6 +21,7 @@ MiniMaxH3MLXPipeline, _adaln_schedule_union, _center_crop_frames, + _configure_metal_memory_limits, _default_metal_wired_limit_gib, _preflight_media_dependencies, _validate_checkpoint_step_ladder, @@ -143,3 +144,26 @@ def fail(*_args, **_kwargs): assert not output.with_suffix(".tmp.mp4").exists() assert not output.with_suffix(".tmp.wav").exists() + + +def test_explicit_wired_limit_uses_wired_api_separately_from_allocator(): + calls = [] + fake = SimpleNamespace( + metal=SimpleNamespace(device_info=lambda: {"memory_size": 36 * 2**30}), + set_memory_limit=lambda size: calls.append(("allocator", size)), + set_wired_limit=lambda size: calls.append(("wired", size)) or 0, + ) + _configure_metal_memory_limits(fake, 27.0) + assert calls == [("allocator", 30 * 2**30), ("wired", 27 * 2**30)] + + +def test_explicit_wired_limit_failure_is_not_silently_ignored(): + def reject(size): + raise ValueError("exceeds system wired limit") + fake = SimpleNamespace(set_wired_limit=reject) + with pytest.raises(ValueError, match="system wired limit"): + _configure_metal_memory_limits(fake, 31.0) + with pytest.raises(ValueError, match="finite and positive"): + _configure_metal_memory_limits(fake, float("nan")) + with pytest.raises(RuntimeError, match="cannot set"): + _configure_metal_memory_limits(SimpleNamespace(), 27.0) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index ca3fe3d11f..f89cc584bd 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -328,3 +328,23 @@ def test_converter_continues_past_mismatched_existing_format(tmp_path, monkeypat converter.main() assert saved == ["int6"] assert (existing / h3.H3_WEIGHTS_FILENAME).read_bytes() == b"existing" + + +@pytest.mark.parametrize("dtype", [mx.float16, mx.bfloat16]) +@pytest.mark.parametrize("exempt", [False, True]) +def test_simd_partial_key_chunks_match_reference(dtype, exempt): + """Nonuniform scores and partial tiles exercise all four 8-key fragments.""" + _require_metal() + geometry = vsa.build_h3_tile_geometry((7, 5), (5, 3, 7), 64) + mx.random.seed(3026) + q, k, value = [mx.random.normal((geometry.total_seq_length, 2, 128)).astype(dtype) + for _ in range(3)] + expected = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="reference") + stats = vsa.MiniMaxH3VSAStats() + actual = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="simd", stats=stats) + mx.eval(actual, expected) + assert stats.impl == "simd" and stats.dense_fallback_reason is None + assert mx.all(mx.isfinite(actual)).item() + # The Metal kernel uses FP32 accumulation with a different reduction order. + np.testing.assert_allclose(np.asarray(actual.astype(mx.float32)), + np.asarray(expected.astype(mx.float32)), atol=.01, rtol=.01) diff --git a/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py new file mode 100644 index 0000000000..7d34e346f4 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU regression for the ModelOpt activation-scale export contract.""" +import importlib.util +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + + +@pytest.fixture +def converter(monkeypatch): + path = Path(__file__).resolve().parents[4] / 'scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py' + spec = importlib.util.spec_from_file_location('h3_modelopt_converter', path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + packed = torch.zeros((2, 8), dtype=torch.uint8) + scales = torch.ones((2, 1), dtype=torch.float8_e4m3fn) + def quantize(*args, **kwargs): + return packed, scales + monkeypatch.setattr(module, '_flashinfer', lambda: ( + SimpleNamespace(layout_128x4=0), None, quantize, lambda value: value)) + return module + + +def convert(converter, input_scale): + return converter.convert_modelopt_linear( + torch.zeros((2, 8), dtype=torch.uint8), + torch.ones((2, 1), dtype=torch.float8_e4m3fn), + torch.tensor(0.25), 'cpu', input_scale=input_scale)[0] + + +def test_modelopt_activation_scale_is_preserved_as_reciprocal(converter): + buffers = convert(converter, torch.tensor(2.0)) + assert buffers['_nvfp4_input_global_sf'].dtype == torch.float32 + assert buffers['_nvfp4_input_global_sf'].shape == () + assert buffers['_nvfp4_input_global_sf'].item() == 0.5 + assert buffers['_nvfp4_alpha'].item() == 0.25 + + +def test_missing_modelopt_activation_scale_keeps_legacy_export(converter): + assert '_nvfp4_input_global_sf' not in convert(converter, None) + + +@pytest.mark.parametrize('value', [0.0, -1.0, float('nan'), float('inf')]) +def test_invalid_modelopt_activation_scale_is_rejected(converter, value): + with pytest.raises(ValueError, match='finite positive scalar'): + convert(converter, torch.tensor(value)) diff --git a/fastvideo/tests/worker/test_ray_distributed_executor.py b/fastvideo/tests/worker/test_ray_distributed_executor.py index 0bf8b4d1e3..ffb3e48421 100644 --- a/fastvideo/tests/worker/test_ray_distributed_executor.py +++ b/fastvideo/tests/worker/test_ray_distributed_executor.py @@ -36,3 +36,17 @@ def test_ray_log_queue_stays_on_the_driver() -> None: executor.clear_log_queue() assert executor._log_queue is None assert "log_queue" in signature(Executor.set_log_queue).parameters + + +def test_ray_carries_h3_performance_switches() -> None: + import fastvideo.envs as envs + from fastvideo.worker.ray_env import get_env_vars_to_copy + + with (envs.FASTVIDEO_H3_VAE_TILE_BATCH.override(8), + envs.FASTVIDEO_NVFP4_MM_BACKEND.override("cutlass"), + envs.FASTVIDEO_VSA_TRITON.override(True)): + copied = get_env_vars_to_copy() + assert {"FASTVIDEO_H3_VAE_TILE_BATCH", "FASTVIDEO_NVFP4_MM_BACKEND", "FASTVIDEO_VSA_TRITON"} <= copied + assert envs.FASTVIDEO_H3_VAE_TILE_BATCH.get() == 8 + assert envs.FASTVIDEO_NVFP4_MM_BACKEND.get() == "cutlass" + assert envs.FASTVIDEO_VSA_TRITON.get() is True diff --git a/fastvideo/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index 4fdcae2c5e..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,7 +22,6 @@ 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.""" @@ -42,14 +42,18 @@ def _log_pipeline_memory(pipeline) -> None: 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])}) + 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, + 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): diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index db96f413c2..4a43adaea2 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -18,6 +18,17 @@ INT8/INT6/INT4 grid, and record ``vsa.capable`` in the manifest. Write VSA checkpoints to a new directory — do not overwrite an existing dense export. +An FP8 transformer with per-channel ``weight_scale`` can also be a source. +The loader dequantizes each FP8 matrix before applying the requested MLX +quantization. This saves download bytes but quantizes twice; compare its clips +with the BF16-sourced export before using it for release. + +``--formats "mxfp8 mxfp4 nvfp4"`` tries native MLX floating-point quantized +storage and matrix multiplication. These formats are experimental and require +operator support from the installed MLX build. This converts BF16 or FP8 source +weights; it does not import CUDA-packed NVFP4 DiT exports. The default formats +remain affine INT8/INT6/INT4. + python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ --out ~/models/FastH3-MLX-vsa \\ @@ -53,8 +64,8 @@ logger = init_logger(__name__) -SUPPORTED_FORMATS = ("int8", "int6", "int4") -DEFAULT_FORMATS = " ".join(SUPPORTED_FORMATS) +SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8", "mxfp4", "nvfp4") +DEFAULT_FORMATS = "int8 int6 int4" def _adaln_cache_timesteps(model_root: str | Path | None = None) -> np.ndarray: @@ -90,6 +101,10 @@ def parse_args() -> argparse.Namespace: help=("retain and quantize transformer_blocks.*.attn.to_gate_compress.weight " "(required for MLX VSA inference; omitted by dense conversion)"), ) + parser.add_argument("--nvfp4-conditioner-root", type=Path, + help="also cache the released packed encoder in native MLX layout, without requantization") + parser.add_argument("--nvfp4-conditioner-out", type=Path, + help="empty encoder-cache output directory; defaults to OUT/nvfp4-encoder") return parser.parse_args() @@ -144,6 +159,15 @@ def main() -> None: if hasattr(mx, "clear_cache"): mx.clear_cache() + if args.nvfp4_conditioner_root is not None: + from fastvideo.mlx_runtime.minimax_h3_conditioner import export_mlx_h3_nvfp4_encoder + + encoder_out = args.nvfp4_conditioner_out or out_base / "nvfp4-encoder" + started = time.perf_counter() + export_mlx_h3_nvfp4_encoder(args.nvfp4_conditioner_root, encoder_out) + print(f"[encoder] cached packed NVFP4 encoder in {time.perf_counter() - started:.1f}s at {encoder_out}", + flush=True) + if __name__ == "__main__": main() diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py index 8b2ff1b448..4f0b3728b7 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -13,10 +13,11 @@ ``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. +``::_weight_global_sf`` (= 1 / ``weight_scale_2``), and +``::_nvfp4_input_global_sf`` (= 1 / ``input_scale``). The calibrated weight +bytes and activation scale are preserved. Dropping the activation scale would +replace its calibrated range with a unit global scale, clipping inputs above +2688. ``--quantize-attention`` additionally quantizes the dense BF16 attention projections (``attn.to_{q,k,v,out}``) of every main block exactly as @@ -105,7 +106,8 @@ def probe(buffers: dict[str, torch.Tensor], reference: torch.Tensor, rows: int = 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]: +def convert_modelopt_linear(weight, scale, scale_2, device, input_scale=None + ) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: """Carry the calibrated bytes over; return (buffers, dequantized weight, scale-byte agreement). The agreement compares the swizzled ModelOpt scales with the scales @@ -128,6 +130,11 @@ def convert_modelopt_linear(weight, scale, scale_2, device) -> tuple[dict[str, t "_nvfp4_alpha": scale_2.clone(), "_weight_global_sf": (1.0 / scale_2).to(torch.bfloat16), } + if input_scale is not None: + value = input_scale.to(dtype=torch.float32) + if value.numel() != 1 or not bool(torch.isfinite(value).all()) or not bool((value > 0).all()): + raise ValueError("ModelOpt input_scale must be a finite positive scalar") + buffers["_nvfp4_input_global_sf"] = value.reshape(()).reciprocal().to(device=device) return buffers, reference, agreement @@ -189,7 +196,9 @@ def main() -> None: else: buffers, reference, agreement = convert_modelopt_linear(tensor(f"{prefix}.weight"), tensor(f"{prefix}.weight_scale"), - tensor(f"{prefix}.weight_scale_2"), device) + tensor(f"{prefix}.weight_scale_2"), device, + input_scale=tensor(f"{prefix}.input_scale") + if f"{prefix}.input_scale" in weight_map else None) agreements.append(agreement) error = probe(buffers, reference) worst = max(worst, error) @@ -233,7 +242,7 @@ def main() -> None: 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, + "static_activation_scales": sum(k.endswith("::_nvfp4_input_global_sf") for k in export), "gate_linears": sum(1 for p in attention if _BLOCK_GATE.match(p)), "min_scale_byte_agreement": round(min(agreements), 4) if agreements else None, "mean_scale_byte_agreement": round(sum(agreements) / len(agreements), 4) if agreements else None,