From e811dac0965d87a098b50ad3a3cde53969e86b0f Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:12:01 -0700 Subject: [PATCH 01/15] [feat]: add resident FastH3 V2 Spark recipe --- docs/getting_started/installation/spark.md | 3 + .../installation/spark_performance.md | 96 +++++++++++++++++-- .../basic/basic_fasth3_spark_v2_nvfp4.yaml | 60 ++++++++++++ 3 files changed, 152 insertions(+), 7 deletions(-) create mode 100644 examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml 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..bea284642b 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,86 @@ 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. +Warm up once, then time at least two `generate_video` calls per prompt in one +process. Record the median wall time, denoise and decode stage times, peak +memory, exact command and commit, and retain every MP4 for review. The V2 +configuration must pass a visual and audio quality check before publication. + +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_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 From fb9c1a290bc160f0a957fedeec4e6bce320add9f Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:29:56 -0700 Subject: [PATCH 02/15] [feat]: benchmark resident FastH3 Spark recipe --- .../installation/spark_performance.md | 19 ++- .../basic/benchmark_fasth3_spark_nvfp4.py | 135 ++++++++++++++++++ 2 files changed, 150 insertions(+), 4 deletions(-) create mode 100644 examples/inference/basic/benchmark_fasth3_spark_nvfp4.py diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index bea284642b..cd616a16ca 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -277,10 +277,21 @@ 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. -Warm up once, then time at least two `generate_video` calls per prompt in one -process. Record the median wall time, denoise and decode stage times, peak -memory, exact command and commit, and retain every MP4 for review. The V2 -configuration must pass a visual and audio quality check before publication. +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. Use the same environment variables as above: + +```bash +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 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..9a82beeb99 --- /dev/null +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -0,0 +1,135 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Time the resident FastH3 V2 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_seconds(result: object) -> dict[str, float]: + 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: float(metrics["execution_time"]) + for name, metrics in stages.items() + if isinstance(metrics, dict) and metrics.get("execution_time") is not 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}") + stages = _stage_seconds(result) + 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": 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() From 476265bd87674bc0292c95406e75ea1b22a9a2b2 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:30:57 -0700 Subject: [PATCH 03/15] [docs]: pin Spark benchmark environment --- docs/getting_started/installation/spark_performance.md | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index cd616a16ca..00c9080d2c 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -279,9 +279,13 @@ the closest legal count above five seconds at 24 fps. For the secondary `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. Use the same environment variables as above: +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 \ From bbb79d9e56e5d53c32bd4915a12a77736ad35c0f Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:01:22 -0700 Subject: [PATCH 04/15] [bugfix]: retain per-stage CUDA peaks in Spark benchmarks --- .../basic/benchmark_fasth3_spark_nvfp4.py | 27 ++++++++++++------- fastvideo/pipelines/stages/base.py | 10 +++++-- 2 files changed, 26 insertions(+), 11 deletions(-) diff --git a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py index 9a82beeb99..f8e3ee7fcd 100644 --- a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Time the resident FastH3 V2 Spark recipe with the release prompts. +"""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 @@ -22,18 +22,24 @@ PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") -def _stage_seconds(result: object) -> dict[str, float]: +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: float(metrics["execution_time"]) - for name, metrics in stages.items() - if isinstance(metrics, dict) and metrics.get("execution_time") is not None - } + 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: @@ -107,7 +113,8 @@ def main() -> None: 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}") - stages = _stage_seconds(result) + metrics = _stage_metrics(result) + stages = _stage_seconds(metrics) row = { "prompt_id": prompt_id, "warmup": warmup, @@ -117,7 +124,9 @@ def main() -> None: "e2e_seconds": round(wall, 3), "denoise_seconds": _stage_total(stages, "denois"), "decode_seconds": _stage_total(stages, "decod"), - "peak_memory_mb": result.peak_memory_mb, + "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), } 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) From 99a03b578da51891700a8354e158dc97468036ef Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:01:22 -0700 Subject: [PATCH 05/15] [feat]: add resident pruned and paired Spark recipes --- .../basic_fasth3_spark_pair_pruned_nvfp4.yaml | 63 +++++++++++++++++++ .../basic_fasth3_spark_pair_v2_nvfp4.yaml | 63 +++++++++++++++++++ .../basic_fasth3_spark_pruned_nvfp4.yaml | 60 ++++++++++++++++++ 3 files changed, 186 insertions(+) create mode 100644 examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml create mode 100644 examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml create mode 100644 examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml 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 From 3f24557a1b870e74608b21fc8433a7d579c034e7 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:09:41 -0700 Subject: [PATCH 06/15] [bugfix]: preserve H3 release helpers after main integration --- fastvideo/models/loader/fsdp_load.py | 2 ++ .../basic/minimax_h3/minimax_h3_pipeline.py | 3 +++ fastvideo/worker/gpu_worker.py | 12 ++++++++---- 3 files changed, 13 insertions(+), 4 deletions(-) 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/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): From e5e5c18b7210592b0b3151fcb6f73cd0ff1e9d87 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:13:49 -0700 Subject: [PATCH 07/15] [feat]: support pruned FastH3 AdaLN in MLX --- fastvideo/mlx_runtime/minimax_h3.py | 40 +++++++++++++++++++---------- 1 file changed, 27 insertions(+), 13 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 58b677bd80..5b5e546cc6 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,6 +1380,10 @@ 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 @@ -1381,7 +1392,7 @@ def assign(key: str, value) -> None: for key, source in shard_arrays.items(): 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,7 +1401,9 @@ 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 + 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) array = _load_array(source, target_dtype) if temb is not None and ".adaln_proj.linear." in key: _, index_str, sub = key.split(".", 2) @@ -1398,7 +1411,7 @@ def assign(key: str, value) -> None: 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 +1441,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"], ) From b8be84c1a8492fb78d6661bd5d20d8f9c743fa86 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:17:23 -0700 Subject: [PATCH 08/15] [feat]: trim MLX H3 conditioner to required layers --- fastvideo/mlx_runtime/minimax_h3_conditioner.py | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index 464250a9c2..b395a9dab1 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -79,7 +79,17 @@ def __init__(self, component_dir: Path): self._header_cache: dict[str, tuple[dict, int]] = {} 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()} + + 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 + + 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(): @@ -294,8 +304,8 @@ def encode_tokens(self, token_ids: list[int]) -> tuple[np.ndarray, np.ndarray]: del rows gc.collect() - 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 From 30754771bd8df65e9795ea0b6f2c5589845b1f48 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:22:43 -0700 Subject: [PATCH 09/15] [feat]: allow FP8 source for MLX H3 conversion --- fastvideo/mlx_runtime/minimax_h3.py | 14 +++++++++++++- .../convert_minimax_h3_mlx.py | 5 +++++ 2 files changed, 18 insertions(+), 1 deletion(-) diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 5b5e546cc6..7cdf1c3dbf 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -1390,6 +1390,8 @@ 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 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.") or key == "adaln_basis.weight"): @@ -1404,7 +1406,17 @@ def assign(key: str, value) -> None: 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) - array = _load_array(source, target_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) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index db96f413c2..1e2813e6ab 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -18,6 +18,11 @@ 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. + python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ --out ~/models/FastH3-MLX-vsa \\ From 892751c6fc806c11d76fd1e3941a9b913fe37c47 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 15:56:41 -0700 Subject: [PATCH 10/15] [docs]: add pruned FastH3 MLX conversion recipe --- docs/cookbook/minimax-h3.md | 5 +++ docs/getting_started/installation/mlx.md | 41 ++++++++++++++++++++++++ 2 files changed, 46 insertions(+) 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..36d20791d9 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -46,6 +46,47 @@ 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/*' + +# The pruned repo's packed NVFP4 encoder is for Blackwell. MLX reads the +# official BF16 encoder through layer 50. The last three shards are unused. +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. + ## Hardware - FastMetal 1.3B and 5B: 16 GB unified memory and up From 80bfd661ad0ec9b299d3e122cbc7f6796ea792aa Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:17:50 -0700 Subject: [PATCH 11/15] [feat]: add resident NVFP4 conditioning for MLX H3 --- docs/getting_started/installation/mlx.md | 45 +++++- .../mlx_runtime/minimax_h3_conditioner.py | 134 +++++++++++++++-- fastvideo/mlx_runtime/minimax_h3_pipeline.py | 95 +++++++++++- .../tests/test_minimax_h3_resident_nvfp4.py | 136 ++++++++++++++++++ 4 files changed, 388 insertions(+), 22 deletions(-) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index 36d20791d9..1f4974b98e 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -58,8 +58,8 @@ hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ --exclude 'text_encoder/*' -# The pruned repo's packed NVFP4 encoder is for Blackwell. MLX reads the -# official BF16 encoder through layer 50. The last three shards are unused. +# 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' \ @@ -87,6 +87,47 @@ 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/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index b395a9dab1..d5fe8b34b8 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 @@ -146,7 +147,7 @@ def close(self) -> 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: @@ -215,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 @@ -248,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): @@ -294,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 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): @@ -378,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_resident_nvfp4.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py new file mode 100644 index 0000000000..1b9f7e2566 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py @@ -0,0 +1,136 @@ +# 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 From 92c46deb2b00eba9483c6b640fb0707fbe3f23e6 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:21:39 -0700 Subject: [PATCH 12/15] [feat]: enable experimental MXFP8 H3 conversion --- scripts/checkpoint_conversion/convert_minimax_h3_mlx.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index 1e2813e6ab..4abcac2805 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -23,6 +23,10 @@ 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 \\ @@ -58,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: From f698e3e0ae49a086751575c911e73abff87e2bb9 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:24:28 -0700 Subject: [PATCH 13/15] [test]: check MXFP8 H3 checkpoint round trip --- .../tests/test_minimax_h3_fp8_checkpoint.py | 27 +++++++++++++++++++ 1 file changed, 27 insertions(+) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py 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 From 068054e36d28f8e3a44c20088df02a6eb818a15d Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:35:38 -0700 Subject: [PATCH 14/15] [bugfix]: omit unused layers from single-shard MLX encoder --- .../mlx_runtime/minimax_h3_conditioner.py | 22 +++++++++---------- .../tests/test_minimax_h3_resident_nvfp4.py | 12 ++++++++++ 2 files changed, 23 insertions(+), 11 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index d5fe8b34b8..f9b1fe5fca 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -78,18 +78,18 @@ 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]] = {} - if index_path.exists(): - weight_map = json.loads(index_path.read_text())["weight_map"] - 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 + 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() if needed(k)} else: single = component_dir / "model.safetensors" @@ -99,7 +99,7 @@ def needed(key: str) -> bool: with open(single, "rb") as handle: (header_len, ) = struct.unpack(" Date: Sat, 3 Oct 2026 21:12:36 -0700 Subject: [PATCH 15/15] [test]: verify rank-16 H3 modulation and cache reload --- .../tests/test_minimax_h3_rank16_adaln.py | 72 +++++++++++++++++++ 1 file changed, 72 insertions(+) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py 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))