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..1f4974b98e 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -46,6 +46,88 @@ 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. + ## 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..00c9080d2c 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,101 @@ 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. Runtime, memory fit, and quality still need a run +on that device. 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. + +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 full H3 VAE. 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. + ## 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/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 58b677bd80..7cdf1c3dbf 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -766,12 +766,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 +914,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 +933,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 +974,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 +1113,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 +1130,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 +1360,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 +1380,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 +1403,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 +1453,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"], ) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index 464250a9c2..f9b1fe5fca 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,96 @@ 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]) + + +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() + + 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): + raw = json.loads((Path(component_dir) / "config.json").read_text()) + quant = raw.get("quantization_config", {}) + expected = { + "quant_method": "nvfp4", + "fmt": "e2m1", + "group_size": 16, + "scale_fmt": "e4m3", + "scale_layout": "128x4", + "activation_scheme": "dynamic" + } + 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.") + # 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) + 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..d55131e49c 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, @@ -328,6 +329,8 @@ 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 @@ -343,6 +346,13 @@ def __init__( except Exception as error: # noqa: BLE001 - best effort on older MLX logger.info("Could not raise the Metal wired limit: %s", error) 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 +424,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 +510,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 +547,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 +697,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 +745,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/tests/test_minimax_h3_fp8_checkpoint.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py new file mode 100644 index 0000000000..a0dc7b3216 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py @@ -0,0 +1,27 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exercise native MXFP8 H3 checkpoint 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, quantize_matrix +from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, save_mlx_h3_checkpoint + + +def test_mxfp8_checkpoint_preserves_quantized_matrix(tmp_path): + spec = MLXQuantizationSpec.from_name('mxfp8') + ensure_quantization_supported(spec) + weight = quantize_matrix(mx.random.normal((64, 64)).astype(mx.bfloat16), spec) + 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 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..6436ea8141 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py @@ -0,0 +1,148 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Native MLX NVFP4 encoder storage and residency regression checks.""" +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, 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]) +def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, global_scale): + 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)) + 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 test_resident_embedding_keeps_bf16_storage(tmp_path): + 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") + 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/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/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/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..4abcac2805 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -18,6 +18,15 @@ 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`` tries native MLX FP8 weight storage and matrix +multiplication. It is experimental and requires operator support from the +installed MLX build; 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 +62,8 @@ logger = init_logger(__name__) -SUPPORTED_FORMATS = ("int8", "int6", "int4") -DEFAULT_FORMATS = " ".join(SUPPORTED_FORMATS) +SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8") +DEFAULT_FORMATS = "int8 int6 int4" def _adaln_cache_timesteps(model_root: str | Path | None = None) -> np.ndarray: