diff --git a/docs/contributing/env_vars.md b/docs/contributing/env_vars.md index 4f47ddcb1f..a834261839 100644 --- a/docs/contributing/env_vars.md +++ b/docs/contributing/env_vars.md @@ -192,6 +192,7 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_FA4` | bool | `0` | attention | The FLASH_ATTN backend uses FlashAttention-4 (flash_attn.cute) instead of FA3 or FA2. | | `FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN` | bool | `0` | attention | MiniMax-H3 dense DiT self-attention uses the FlashAttention-4 packed-varlen entry point. This changes the floating-point reduction order, so it is an inference-only opt-in. | | `FASTVIDEO_VSA_SM100A` | bool | `0` | attention | VIDEO_SPARSE_ATTN_H3 sends no-grad tile-64 forwards to the data-center Blackwell (sm_100a) kernel. fastvideo-kernel reads the same variable with the same rule. | +| `FASTVIDEO_VSA_TRITON` | bool | `0` | attention | Force the Triton MiniMax-H3 sparse attention kernel. fastvideo-kernel reads the same variable. | | `FASTVIDEO_NVFP4_FA4` | bool | `0` | attention | FlashAttention-4 quantizes Q and K to NVFP4. An explicit nvfp4_fa4 attention implementation argument takes precedence. | | `FASTVIDEO_DISABLE_ATTENTION_COMPILE` | bool | `1` | attention | Keep attention forward out of torch.compile graphs (torch.compiler.disable). Set it to 0 to let attention constructed under that setting be traced. Setting it explicitly to true also blocks regional compile. | | `FASTVIDEO_MLX_WINDOW` | int | `0` | attention | MLX FastWan windowed attention size in tokens. 0 uses full attention. | @@ -201,6 +202,8 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_VAE_PARALLEL_ENCODE` | bool | `0` | performance | MiniMax-H3 reference-video VAE encode splits its temporal chunks across the sequence-parallel ranks. Same as FastVideoArgs.vae_parallel_encode=True. | | `FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY` | str | unset | performance | Collective that moves chunks in parallel VAE decode: gather (used when unset) or all_gather. | | `FASTVIDEO_MINIMAX_H3_FUSIONS` | str | `""` | performance | MiniMax-H3 inference-only Triton fusions: all, 1, or a comma-separated subset of modulate,qknorm_rope,swiglu. Empty, 0, or none keeps the eager implementation. | +| `FASTVIDEO_H3_VAE_TILE_BATCH` | int | `1` | performance | Spatial tiles per MiniMax-H3 light-VAE decoder call. Values below 1 use one tile. | +| `FASTVIDEO_NVFP4_MM_BACKEND` | str | `auto` | performance | FlashInfer NVFP4 matrix multiplication backend: auto, cutlass, cudnn, trtllm, or b12x. | | `FASTVIDEO_FSDP2_AUTOWRAP` | bool | `0` | performance | FSDP2 shards modules by parameter count instead of the model's shard conditions. Not supported by self-forcing distillation. | | `FASTVIDEO_FSDP2_MIN_PARAMS` | int | `10000000` | performance | Minimum parameter count of a module that FASTVIDEO_FSDP2_AUTOWRAP shards. | | `FASTVIDEO_MLX_COMPILE` | bool | `0` | performance | Compile the MLX DiT forward with mx.compile. | diff --git a/docs/cookbook/minimax-h3.md b/docs/cookbook/minimax-h3.md index 93c8df76c7..b82762f654 100644 --- a/docs/cookbook/minimax-h3.md +++ b/docs/cookbook/minimax-h3.md @@ -11,6 +11,11 @@ a full model, not a demo. **V2** is the eight-step checkpoint. More forwards is why V2 is the higher-quality FastH3. The V2 schedule contract is in [FastH3 distilled checkpoint schedules](../inference/fasth3-distilled.md). +The 42-block pruned checkpoint has an [MLX INT8/INT6 conversion and +eight-forward T2VA command](../getting_started/installation/mlx.md#pruned-eight-forward-checkpoint). +It reads `fastvideo_inference.json` for the trained schedule. The command +uses native 832x480 resolution and all requested frames. +
None: gc.collect() -_DTYPES = {"F32": np.float32, "F16": np.float16, "I64": np.int64, "I32": np.int32} +_DTYPES = {"F32": np.float32, "F16": np.float16, "I64": np.int64, "I32": np.int32, "U8": np.uint8} def _read_bf16_words(path: str, key: str, header: dict, data_start: int) -> np.ndarray: @@ -205,8 +216,20 @@ def _rms_norm(x, weight, eps: float): return x / mx.sqrt(mx.mean(x * x, axis=-1, keepdims=True) + eps) * weight +@dataclass(frozen=True) +class NVFP4Matrix: + """MLX row-major E2M1/E4M3 weights and the export's inverse global scale.""" + + weight: mx.array + scales: mx.array + global_scale: float + + def matmul(self, x): + return mx.quantized_matmul(x, self.weight, self.scales, mode="nvfp4") / self.global_scale + + def _linear(x, weight, bias=None): - y = x @ weight.T + y = weight.matmul(x) if isinstance(weight, NVFP4Matrix) else x @ weight.T if bias is not None: y = y + bias return y @@ -238,7 +261,7 @@ class StreamedMiniMaxH3TextConditioner: def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): self.component_dir = Path(component_dir) self.config = ConditionerConfig.from_config_json(self.component_dir / "config.json") - self.index = _ShardIndex(self.component_dir) + self.index: Any = _ShardIndex(self.component_dir) self.tokenizer = self._load_tokenizer(tokenizer_dir) def _load_tokenizer(self, tokenizer_dir: str | Path | None): @@ -284,28 +307,31 @@ def encode_tokens(self, token_ids: list[int]) -> tuple[np.ndarray, np.ndarray]: ]) cos, sin = _mrope_cos_sin(positions, cfg) - # Embedding rows gathered individually; the (151936, 5120) table is - # never fully materialized. - rows = [] - for token in token_ids: - key = "model.language_model.embed_tokens.weight" - rows.append(self.index.get_row(key, token)) - hidden = mx.array(np.stack(rows).astype(np.float32)) - del rows - gc.collect() + hidden = self._embed_tokens(token_ids) - if cfg.num_layers <= TEXT_ENCODER_LAYER: - raise ValueError(f"Conditioner needs > {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") + if cfg.num_layers < TEXT_ENCODER_LAYER: + raise ValueError(f"Conditioner needs at least {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") for layer in range(TEXT_ENCODER_LAYER): hidden = self._decoder_layer(layer, hidden, cos, sin) - # Per-layer sync: without this the whole 50-layer graph accumulates - # and the machine runs out of memory (same failure mode as the DiT). + # Per-layer sync keeps the 50-layer activation graph bounded. mx.eval(hidden) gc.collect() tags = np.full((seq_len, ), 1, dtype=np.int64) # MINIMAX_H3_TEXT_TAG return np.asarray(hidden).astype(np.float32), tags + def _embed_tokens(self, token_ids: list[int]): + # Embedding rows gathered individually; the full table is never + # materialized by the BF16 streaming path. + rows = [] + for token in token_ids: + key = "model.language_model.embed_tokens.weight" + rows.append(self.index.get_row(key, token)) + hidden = mx.array(np.stack(rows).astype(np.float32)) + del rows + gc.collect() + return hidden + # -- layers ---------------------------------------------------------- def _decoder_layer(self, index: int, hidden, cos, sin): @@ -368,6 +394,171 @@ def close(self) -> None: self.index.close() +def unswizzle_nvfp4_scales(scale: np.ndarray, rows: int, cols: int) -> np.ndarray: + """FlashInfer 128x4 scale bytes -> MLX row-major group-16 scale bytes.""" + pad_rows, pad_cols = -(-rows // 128) * 128, -(-cols // 4) * 4 + if scale.size != pad_rows * pad_cols: + raise ValueError(f"NVFP4 scales need {pad_rows * pad_cols} bytes, got {scale.size}.") + tiles = scale.reshape(pad_rows // 128, pad_cols // 4, 32, 4, 4) + return np.ascontiguousarray(tiles.transpose(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols]) + + +MLX_NVFP4_ENCODER_MANIFEST = "mlx_h3_nvfp4_encoder.json" + + +def _read_nvfp4_encoder_config(component_dir: str | Path) -> dict[str, Any]: + raw = json.loads((Path(component_dir) / "config.json").read_text()) + expected = { + "quant_method": "nvfp4", + "fmt": "e2m1", + "group_size": 16, + "scale_fmt": "e4m3", + "scale_layout": "128x4", + "activation_scheme": "dynamic", + } + quant = raw.get("quantization_config", {}) + if any(quant.get(key) != value for key, value in expected.items()): + raise ValueError("MLX NVFP4 conditioning requires the FastVideo group-16, 128x4 encoder export.") + return raw + + +def export_mlx_h3_nvfp4_encoder(component_dir: str | Path, output_dir: str | Path) -> Path: + """Cache the released packed encoder in MLX layout, without requantization. + + Keep original packed nibbles, row-major scale bytes, global scales and + embedding/norm values. Later loads skip CPU scale unswizzling and staging. + """ + raw = _read_nvfp4_encoder_config(component_dir) + output_dir = Path(output_dir) + if output_dir.exists() and any(output_dir.iterdir()): + raise FileExistsError(f"Encoder cache output must be empty: {output_dir}") + output_dir.mkdir(parents=True, exist_ok=True) + index = _ResidentNVFP4Index(_ShardIndex(Path(component_dir))) + try: + arrays = {} + matrices = {} + dense_keys = [] + for key, value in index.weights.items(): + if isinstance(value, NVFP4Matrix): + arrays[key] = value.weight + arrays[key + ".scales"] = value.scales + matrices[key] = {"global_scale": value.global_scale} + else: + arrays[key] = value + dense_keys.append(key) + mx.save_safetensors(str(output_dir / "model.safetensors"), arrays) + (output_dir / "config.json").write_text(json.dumps(raw, indent=2) + "\n") + manifest = { + "format_version": 1, + "language_layers": TEXT_ENCODER_LAYER, + "matrices": matrices, + "dense_keys": dense_keys, + "source_dir": str(Path(component_dir).resolve()) + } + (output_dir / MLX_NVFP4_ENCODER_MANIFEST).write_text(json.dumps(manifest, indent=2) + "\n") + finally: + index.close() + return output_dir + + +class _ResidentNVFP4Index: + """Load the released 50-layer encoder without expanding packed matrices.""" + + def __init__(self, source: _ShardIndex): + self.weights: dict[str, mx.array | NVFP4Matrix] = {} + for key in sorted(source.key_to_shard): + if key.endswith(".weight_packed"): + prefix = key.removesuffix(".weight_packed") + packed = np.array(source.get(key), copy=True) + if packed.dtype != np.uint8 or packed.ndim != 2 or packed.shape[1] % 4: + raise ValueError(f"Invalid packed NVFP4 matrix {key}: {packed.shape}, {packed.dtype}") + rows, cols = packed.shape[0], packed.shape[1] * 2 + if cols % 16: + raise ValueError(f"NVFP4 input width must be divisible by 16: {key}") + scales = unswizzle_nvfp4_scales(source.get(prefix + ".weight_scale"), rows, cols // 16) + global_scale = float(source.get(prefix + ".weight_global_scale").reshape(-1)[0]) + if not np.isfinite(global_scale) or global_scale <= 0: + raise ValueError(f"Invalid NVFP4 global scale for {prefix}: {global_scale}") + weight = mx.array(packed).view(mx.uint32) + scale_bytes = mx.array(scales) + if not bool(mx.all(mx.isfinite(mx.from_fp8(scale_bytes, dtype=mx.float32)))): + raise ValueError(f"Non-finite NVFP4 block scales for {prefix}") + mx.eval(weight, scale_bytes) + self.weights[prefix + ".weight"] = NVFP4Matrix(weight, scale_bytes, global_scale) + elif key.endswith(".weight"): + if key == "model.language_model.embed_tokens.weight": + shard = source.key_to_shard[key] + header, data_start = source._cache_header(shard) + if header[key]["dtype"] == "BF16": + value = mx.array(_read_bf16_words(shard, key, header, data_start)).view(mx.bfloat16) + else: + value = mx.array(source.get(key)) + else: + value = source.get_mlx(key) + mx.eval(value) + self.weights[key] = value + source.close() + + @classmethod + def from_mlx_checkpoint(cls, component_dir: str | Path): + component_dir = Path(component_dir) + manifest = json.loads((component_dir / MLX_NVFP4_ENCODER_MANIFEST).read_text()) + if manifest.get("format_version") != 1 or manifest.get("language_layers") != TEXT_ENCODER_LAYER: + raise ValueError("Unsupported native MLX NVFP4 encoder cache") + arrays = mx.load(str(component_dir / "model.safetensors")) + matrices = manifest["matrices"] + expected = set(manifest["dense_keys"]) | set(matrices) | {key + ".scales" for key in matrices} + if set(arrays) != expected: + raise ValueError("Native MLX encoder arrays do not match the manifest") + index = cls.__new__(cls) + index.weights = {key: arrays[key] for key in manifest["dense_keys"]} + for key, info in matrices.items(): + weight, scales = arrays[key], arrays[key + ".scales"] + factor = float(info["global_scale"]) + if (weight.ndim != 2 or weight.dtype != mx.uint32 or scales.dtype != mx.uint8 + or scales.shape != (weight.shape[0], weight.shape[1] // 2) or weight.shape[1] % 2 + or not np.isfinite(factor) or factor <= 0): + raise ValueError(f"Invalid native MLX NVFP4 matrix: {key}") + index.weights[key] = NVFP4Matrix(weight, scales, factor) + mx.eval(list(arrays.values())) + return index + + def get_mlx(self, key: str): + return self.weights[key] + + def close(self) -> None: + self.weights.clear() + gc.collect() + + +class ResidentNVFP4MiniMaxH3TextConditioner(StreamedMiniMaxH3TextConditioner): + """Released NVFP4 encoder weights with floating-point MLX activations. + + The packed weights and embedding table stay resident. CUDA quantizes + activations to FP4; this path keeps FP32 activations, so hidden states are + not expected to be bit-exact with the CUDA encoder. + """ + + def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): + _read_nvfp4_encoder_config(component_dir) + # Fail on an older MLX before reading the encoder's large shards. + try: + packed, scales = mx.quantize(mx.ones((1, 64)), mode="nvfp4") + mx.eval(mx.quantized_matmul(mx.ones((1, 64)), packed, scales, mode="nvfp4")) + except (ValueError, RuntimeError) as error: + raise RuntimeError("Native NVFP4 conditioning requires an MLX build with nvfp4 matmul support.") from error + super().__init__(component_dir, tokenizer_dir) + if (Path(component_dir) / MLX_NVFP4_ENCODER_MANIFEST).exists(): + self.index.close() + self.index = _ResidentNVFP4Index.from_mlx_checkpoint(component_dir) + else: + self.index = _ResidentNVFP4Index(self.index) + + def _embed_tokens(self, token_ids: list[int]): + table = self.index.get_mlx("model.language_model.embed_tokens.weight") + return table[mx.array(token_ids, dtype=mx.int32)].astype(mx.float32) + + def _apply_mrope(q_or_k, cos, sin): """q_or_k: (S, H, D); cos/sin: (S, 1, D).""" half = q_or_k.shape[-1] // 2 diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index c6ca112f32..be898f3bb8 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -49,6 +49,7 @@ MINIMAX_H3_MIN_DURATION, MINIMAX_H3_VIDEO_SHIFT, MiniMaxH3SchedulerState, + _eval_value, adaln_timestep_union, align_num_frames, audio_latent_num_frames, @@ -231,7 +232,7 @@ def _cleanup_mlx() -> None: def _default_metal_wired_limit_gib(mx) -> float: - """Keep the default below both physical memory and the tested 30 GiB cap.""" + """Legacy helper for allocator capacity, not the wired-residency setting.""" metal = getattr(mx, "metal", None) if metal is None: return 30.0 @@ -244,6 +245,30 @@ def _default_metal_wired_limit_gib(mx) -> float: return min(30.0, 0.84 * total_bytes / 2**30) +def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> None: + """Keep allocator capacity separate from explicitly requested wired residency.""" + set_memory = getattr(mx, "set_memory_limit", None) + if set_memory is None and hasattr(mx, "metal"): + set_memory = getattr(mx.metal, "set_memory_limit", None) + if set_memory is not None: + try: + set_memory(int(_default_metal_wired_limit_gib(mx) * 2**30)) + except Exception as error: # noqa: BLE001 - older MLX best effort + logger.info("Could not set the Metal allocation limit: %s", error) + if wired_limit_gib is None: + return + if not math.isfinite(wired_limit_gib) or wired_limit_gib <= 0: + raise ValueError("metal_wired_limit_gib must be finite and positive") + set_wired = getattr(mx, "set_wired_limit", None) + if set_wired is None and hasattr(mx, "metal"): + set_wired = getattr(mx.metal, "set_wired_limit", None) + if set_wired is None: + raise RuntimeError("This MLX build cannot set the requested wired-memory limit") + # Explicit requests must succeed; do not silently benchmark an unwired model. + previous = set_wired(int(wired_limit_gib * 2**30)) + logger.info("MLX wired limit %.2f GiB (previous %.2f GiB)", wired_limit_gib, previous / 2**30) + + MINIMAX_H3_PROMPT_CACHE_VERSION = "v2-attention-layout" @@ -328,21 +353,20 @@ def __init__( video_decode_backend: str = "h3-vae", taeh3_checkpoint: str | Path | None = None, taeh3_chunk_size: int = 5, + conditioner_mode: str = "auto", + resident: bool = False, ) -> None: import mlx.core as mx - set_limit = getattr(mx, "set_memory_limit", None) - if set_limit is None and hasattr(mx, "metal"): - set_limit = getattr(mx.metal, "set_memory_limit", None) - if set_limit is not None: - # Keep large resident models inside a predictable wired budget. - try: - if metal_wired_limit_gib is None: - metal_wired_limit_gib = _default_metal_wired_limit_gib(mx) - set_limit(int(metal_wired_limit_gib * 2**30)) - except Exception as error: # noqa: BLE001 - best effort on older MLX - logger.info("Could not raise the Metal wired limit: %s", error) + _configure_metal_memory_limits(mx, metal_wired_limit_gib) self.model_root = Path(model_root) + if conditioner_mode not in ("auto", "streamed", "nvfp4"): + raise ValueError(f"Unknown H3 conditioner mode: {conditioner_mode}") + if resident and video_decode_backend != "h3-vae": + raise ValueError("Resident H3 generation requires the H3 video VAE.") + self.conditioner_mode = conditioner_mode + self.resident = resident + self._resident_components: dict[str, Any] = {} self.dit_checkpoint = Path(mlx_dit_checkpoint) self.vae_dtype = vae_dtype if video_decode_backend not in ("h3-vae", "taeh3"): @@ -414,20 +438,70 @@ def resolve_geometry( # -- phase 1: conditioning ------------------------------------------- + def prepare_resident(self) -> None: + """Load the encoder, DiT, and both decoders once, before timed requests.""" + if not self.resident or self._resident_components: + return + from fastvideo.mlx_runtime.minimax_h3_conditioner import ResidentNVFP4MiniMaxH3TextConditioner + from fastvideo.mlx_runtime.minimax_h3_audio_vae import mlx_h3_audio_vae_from_dir + from fastvideo.mlx_runtime.minimax_h3_video_vae import mlx_h3_video_vae_from_dir + + import mlx.core as mx + + try: + conditioner = self._load_conditioner() + if not isinstance(conditioner, ResidentNVFP4MiniMaxH3TextConditioner): + conditioner.close() + raise ValueError("All-resident generation requires the packed NVFP4 text encoder.") + self._resident_components["conditioner"] = conditioner + dit = load_mlx_h3_checkpoint(self.dit_checkpoint) + self._resident_components["dit"] = dit + for group in [dit.weights, *dit.blocks, *dit.refiner]: + for value in group.values(): + _eval_value(value) + cache = dit._adaln_cache + if cache is not None: + mx.eval(cache.block_tables, cache.norm_out_shift, cache.norm_out_scale) + self._resident_components["video_vae"] = mlx_h3_video_vae_from_dir(self.model_root / "vae", + include_encoder=False, + storage_dtype=self.vae_dtype) + self._resident_components["audio_vae"] = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", + include_encoder=False) + mx.eval(list(self._resident_components["audio_vae"].weights.values())) + logger.info("H3 components resident: %.2f GiB active MLX memory", mx.get_active_memory() / 2**30) + except Exception: + self.close() + raise + + def close(self) -> None: + conditioner = self._resident_components.get("conditioner") + if conditioner is not None: + conditioner.close() + self._resident_components.clear() + _cleanup_mlx() + def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: """Returns (hidden states (S, hidden), token tags). Uses the cache or the streamed conditioner.""" cache_key = None if self.prompt_cache_dir is not None: - cache_key = prompt_cache_path(self.prompt_cache_dir, self.model_root, prompt) + identity = (f"{self.model_root}::conditioner=" + f"{getattr(self, 'conditioner_dir', self.model_root / 'text_encoder')}::" + f"{getattr(self, 'conditioner_mode', 'auto')}") + cache_key = prompt_cache_path(self.prompt_cache_dir, identity, prompt) if cache_key.exists(): data = np.load(cache_key) logger.info("Loaded prompt embeddings from cache %s", cache_key) return data["hidden_states"], data["token_tags"] - conditioner = self._load_conditioner() + if getattr(self, "resident", False): + self.prepare_resident() + conditioner = self._resident_components["conditioner"] + else: + conditioner = self._load_conditioner() hidden, tags = conditioner.encode_prompt(prompt) - conditioner.close() + if not getattr(self, "resident", False): + conditioner.close() _cleanup_mlx() if cache_key is not None: cache_key.parent.mkdir(parents=True, exist_ok=True) @@ -450,8 +524,17 @@ def has_conditioner_weights(self) -> bool: return marker.exists() or single.exists() def _load_conditioner(self): - from fastvideo.mlx_runtime.minimax_h3_conditioner import StreamedMiniMaxH3TextConditioner + from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + ResidentNVFP4MiniMaxH3TextConditioner, + StreamedMiniMaxH3TextConditioner, + ) + config = json.loads((self.conditioner_dir / "config.json").read_text()) + packed = config.get("quantization_config", {}).get("quant_method") == "nvfp4" + if self.conditioner_mode == "nvfp4" or (self.conditioner_mode == "auto" and packed): + return ResidentNVFP4MiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + if packed: + raise ValueError("The streamed conditioner requires BF16 weights; use conditioner_mode='nvfp4'.") return StreamedMiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) # -- phase 2: denoise -------------------------------------------------- @@ -478,6 +561,10 @@ def denoise( geometry = self.resolve_geometry(height, width, num_frames, enforce_duration=audio_num_frames is None) audio_frames = geometry["num_frames"] if audio_num_frames is None else align_num_frames(audio_num_frames) + if dit is None and getattr(self, "resident", False): + _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) + self.prepare_resident() + dit = self._resident_components["dit"] owned_dit = dit is None if owned_dit: _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) @@ -624,7 +711,13 @@ def decode_video(self, raise RuntimeError(f"TAEH3 produced unexpected frame shape: {frames.shape}") _cleanup_mlx() return frames - vae = mlx_h3_video_vae_from_dir(self.model_root / "vae", include_encoder=False, storage_dtype=self.vae_dtype) + if getattr(self, "resident", False): + self.prepare_resident() + vae = self._resident_components["video_vae"] + else: + vae = mlx_h3_video_vae_from_dir(self.model_root / "vae", + include_encoder=False, + storage_dtype=self.vae_dtype) expected_height = height // vae.spatial_compression_ratio expected_width = width // vae.spatial_compression_ratio if (geometry["latent_height"], geometry["latent_width"]) != (expected_height, expected_width): @@ -666,7 +759,11 @@ def decode_audio(self, audio_rows: np.ndarray, *, num_frames: int) -> np.ndarray num_audio_latents = audio_latent_num_frames(align_num_frames(num_frames)) latents = unpack_audio_tokens(audio_rows, num_audio_latents) - vae = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", include_encoder=False) + if getattr(self, "resident", False): + self.prepare_resident() + vae = self._resident_components["audio_vae"] + else: + vae = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", include_encoder=False) z = vae.denormalize_latents(mx.array(latents)) waveform = np.asarray(vae.decode(z))[:, 0, :] # (B, 1, S) -> (B, S) del vae, z diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa.py b/fastvideo/mlx_runtime/minimax_h3_vsa.py index 3c49969e5c..48ca15fa32 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa.py @@ -544,11 +544,11 @@ def _key_valid_mask(block_idx, variable_block_sizes, tile_elems: int): return offsets[None, None, None, :] < selected_sizes[:, :, :, None] -_REFERENCE_GATHER_TARGET_BYTES = 2 * 1024**3 +_REFERENCE_GATHER_TARGET_BYTES = 256 * 1024**2 def _reference_gather_query_chunk(heads: int, dim: int, k_sel: int, tile_elems: int, n_q: int) -> int: - """Batch as many query tiles as fit in ~2 GiB of gathered BF16 K/V.""" + """Bound gathered BF16 K/V to 256 MiB, leaving space for resident weights.""" bytes_per_query = 4 * heads * max(k_sel, 1) * tile_elems * dim chunk = min(n_q, max(1, _REFERENCE_GATHER_TARGET_BYTES // max(bytes_per_query, 1))) return int(chunk) diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index bb21b1912a..14068f14de 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -37,16 +37,18 @@ """ # One threadgroup = one (head, video query tile). 8 SIMD-groups x 32 = 256 -# threads cover 64 query rows. Q is half smem (16 KiB). K/V stage 8 keys. +# threads cover 64 query rows. K/V stage 32 keys, with 28.25 KiB total smem. +# Updating online softmax once per 32 keys reduces accumulator rescaling barriers. +# Four SIMD lanes cooperate per query row in the softmax reduction. _SIMD_SOURCE = """ const int TILE = 64; const int D = 128; const int SG = 32; const int N_SG = 8; const int ROWS = 8; - const int KCHUNK = 8; - threadgroup float kvsmem[8 * 128]; - threadgroup float score_smem[8 * 64]; + const int KCHUNK = 32; + threadgroup float kvsmem[KCHUNK * D]; + threadgroup float score_smem[N_SG * ROWS * KCHUNK]; threadgroup float scale_tmp[8 * 64]; threadgroup float qtile[8 * 64]; threadgroup float row_alpha[8 * 8]; @@ -65,7 +67,7 @@ int qt = n_prefix + (int)q_tile; int q_valid = active ? vbs[qt] : 0; int q_base_tile = (((int)head * S) + qt * TILE) * D; - threadgroup float *sg_scores = score_smem + sid * 64; + threadgroup float *sg_scores = score_smem + sid * ROWS * KCHUNK; threadgroup float *sg_tmp = scale_tmp + sid * 64; threadgroup float *sg_qtile = qtile + sid * 64; threadgroup float *sg_alpha = row_alpha + sid * 8; @@ -117,34 +119,38 @@ } threadgroup_barrier(mem_flags::mem_threadgroup); - thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 kmat; - simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kk * 8), D, ulong2(0, 0), true); - simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); + for (int kc = 0; kc < KCHUNK; kc += 8) { + thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix (0.0f); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 kmat; + simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D, ulong2(0, 0), true); + simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); + } + simdgroup_store(smat, sg_scores + kc, KCHUNK); } - simdgroup_store(smat, sg_scores, KCHUNK); simdgroup_barrier(mem_flags::mem_threadgroup); - float scores[8]; + float scores[KCHUNK / 4]; float cmax = -3.402823466e+38f; - if (lane < (uint)ROWS) { - int grow = qrow0 + (int)lane; - for (int t = 0; t < KCHUNK; t++) { - int gtok = j0 + t; - float sc = -3.402823466e+38f; - if (grow < q_valid && gtok < k_valid && gtok < TILE) { - sc = sg_scores[(int)lane * KCHUNK + t] * scale; - } - scores[t] = sc; - cmax = metal::max(cmax, sc); + int row = (int)lane / 4; + int col_lane = (int)lane % 4; + int grow = qrow0 + row; + for (int t = col_lane; t < KCHUNK; t += 4) { + int gtok = j0 + t; + float sc = -3.402823466e+38f; + if (grow < q_valid && gtok < k_valid && gtok < TILE) { + sc = sg_scores[row * KCHUNK + t] * scale; } - float m_new = metal::max(row_m, cmax); - float alpha = metal::exp(row_m - m_new); - row_lse *= alpha; - row_m = m_new; - sg_alpha[(int)lane] = alpha; + scores[t / 4] = sc; + cmax = metal::max(cmax, sc); } + cmax = metal::max(cmax, simd_shuffle_xor(cmax, 1)); + cmax = metal::max(cmax, simd_shuffle_xor(cmax, 2)); + float m_new = metal::max(row_m, cmax); + float alpha = metal::exp(row_m - m_new); + row_lse *= alpha; + row_m = m_new; + if (col_lane == 0) sg_alpha[row] = alpha; simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { scale_rows_simd8x8(acc[kk], sg_alpha, sg_tmp, lane); @@ -164,33 +170,32 @@ threadgroup_barrier(mem_flags::mem_threadgroup); float local = 0.0f; - if (lane < (uint)ROWS) { - int grow = qrow0 + (int)lane; - for (int t = 0; t < KCHUNK; t++) { - float w = 0.0f; - if (grow < q_valid) { - w = metal::exp(scores[t] - row_m); - } - sg_scores[(int)lane * KCHUNK + t] = w; - local += w; - } - row_lse += local; + for (int t = col_lane; t < KCHUNK; t += 4) { + float w = 0.0f; + if (grow < q_valid) w = metal::exp(scores[t / 4] - row_m); + sg_scores[row * KCHUNK + t] = w; + local += w; } + local += simd_shuffle_xor(local, 1); + local += simd_shuffle_xor(local, 2); + row_lse += local; simdgroup_barrier(mem_flags::mem_threadgroup); - thread simdgroup_float8x8 pmat; - simdgroup_load(pmat, sg_scores, KCHUNK); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 vmat; - simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kk * 8), D); - simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); + for (int kc = 0; kc < KCHUNK; kc += 8) { + thread simdgroup_float8x8 pmat; + simdgroup_load(pmat, sg_scores + kc, KCHUNK); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 vmat; + simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D); + simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); + } } threadgroup_barrier(mem_flags::mem_threadgroup); } } - if (lane < (uint)ROWS) { - sg_alpha[(int)lane] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; + if (lane % 4 == 0) { + sg_alpha[(int)lane / 4] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; } simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { @@ -211,6 +216,7 @@ } simdgroup_barrier(mem_flags::mem_threadgroup); } + """ diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py new file mode 100644 index 0000000000..caac1541e8 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exercise native floating-point quantized H3 storage on supported MLX builds.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +from fastvideo.mlx_runtime.fastwan import MLXQuantizationSpec, ensure_quantization_supported, linear +from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, quantize_matrix, save_mlx_h3_checkpoint + + +@pytest.mark.parametrize('mode', ['mxfp8', 'mxfp4', 'nvfp4']) +def test_float_quantized_checkpoint_preserves_matrix(tmp_path, mode): + spec = MLXQuantizationSpec.from_name(mode) + ensure_quantization_supported(spec) + dense = (mx.random.normal((64, 64)) * 0.001).astype(mx.bfloat16) + weight = quantize_matrix(dense, spec) + restored = mx.dequantize(weight.weight, weight.scales, mode=mode).astype(mx.float32) * weight.global_scale + relative_error = mx.sqrt(mx.sum((restored - dense.astype(mx.float32))**2) / mx.sum(dense.astype(mx.float32)**2)) + assert float(relative_error.item()) < 0.15 + x = mx.random.normal((3, 64)).astype(mx.bfloat16) + config = dict(hidden_size=64, num_attention_heads=1, attention_head_dim=64, ffn_dim=128, + in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=64, + freq_dim=64, time_embed_dim=64, rope_freq_dim=4, rope_theta=10000., + norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5) + dit = MLXMiniMaxH3DiT({'test.weight': weight}, [], [], config) + save_mlx_h3_checkpoint(dit, tmp_path) + loaded = load_mlx_h3_checkpoint(tmp_path) + actual = linear(x, loaded.weights['test.weight']).astype(mx.float32) + expected = linear(x, weight).astype(mx.float32) + np.testing.assert_array_equal(np.array(actual), np.array(expected)) + assert loaded.weights['test.weight'].biases is None + assert loaded.weights['test.weight'].spec == spec + assert loaded.weights['test.weight'].global_scale == weight.global_scale diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py new file mode 100644 index 0000000000..cd954981b6 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py @@ -0,0 +1,72 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Check the pruned model's shared rank-16 modulation against NumPy math.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +from fastvideo.mlx_runtime.fastwan import timestep_embedding +from fastvideo.mlx_runtime.minimax_h3 import ( + MLXMiniMaxH3DiT, + MiniMaxH3SchedulerState, + adaln_timestep_union, + load_mlx_h3_checkpoint, + save_mlx_h3_checkpoint, +) + + +def test_rank16_cache_has_one_silu_before_shared_basis(tmp_path): + rng = np.random.default_rng(2026) + hidden, rank = 32, 16 + + def array(shape): + return mx.array(rng.normal(0, 0.1, shape).astype(np.float32)) + + config = dict(hidden_size=hidden, num_attention_heads=1, attention_head_dim=hidden, ffn_dim=64, + in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=hidden, + freq_dim=hidden, time_embed_dim=hidden, rope_freq_dim=4, rope_theta=10000., + norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5, adaln_rank=rank, num_layers=42) + weights = { + 'time_embedder.linear_1.weight': array((hidden, hidden)), + 'time_embedder.linear_1.bias': array((hidden,)), + 'time_embedder.linear_2.weight': array((hidden, hidden)), + 'time_embedder.linear_2.bias': array((hidden,)), + 'adaln_basis.weight': array((rank, hidden)), + 'norm_out.linear.weight': array((2 * hidden, rank)), + 'norm_out.linear.bias': array((2 * hidden,)), + } + blocks = [{'attn.to_q.weight': array((hidden, hidden)), + 'adaln_proj.linear.weight': array((18 * hidden, rank)), + 'adaln_proj.linear.bias': array((18 * hidden,))} for _ in range(42)] + dit = MLXMiniMaxH3DiT(weights, blocks, [], config) + rungs = [999, 874, 749, 624, 500, 375, 250, 125] + timesteps = adaln_timestep_union(MiniMaxH3SchedulerState.from_dmd_steps(10, rungs), + MiniMaxH3SchedulerState.from_dmd_steps(3, rungs)) + + def project(x, weight, bias=None): + result = x @ np.array(weight).T + return result if bias is None else result + np.array(bias) + + def silu(x): + return x / (1 + np.exp(-x)) + + features = np.array(timestep_embedding(mx.array(timesteps), hidden)) + first = project(features, weights['time_embedder.linear_1.weight'], weights['time_embedder.linear_1.bias']) + second = project(silu(first), weights['time_embedder.linear_2.weight'], weights['time_embedder.linear_2.bias']) + expected_basis = project(silu(second), weights['adaln_basis.weight']) + np.testing.assert_allclose(np.array(dit.compute_temb(mx.array(timesteps))), expected_basis, atol=1e-6) + expected_blocks = [project(expected_basis, block['adaln_proj.linear.weight'], + block['adaln_proj.linear.bias']).reshape(-1, 6 * hidden) for block in blocks] + expected_out = project(expected_basis, weights['norm_out.linear.weight'], weights['norm_out.linear.bias']) + cache = dit.precompute_adaln(timesteps) + for tables, expected in zip(cache.block_tables, expected_blocks, strict=True): + np.testing.assert_allclose(np.concatenate([np.array(t) for t in tables], axis=-1), expected, atol=1e-6) + np.testing.assert_allclose(np.array(cache.norm_out_shift), expected_out[:, :hidden], atol=1e-6) + np.testing.assert_allclose(np.array(cache.norm_out_scale), expected_out[:, hidden:], atol=1e-6) + assert all(block['adaln_proj.linear.weight'] is None for block in blocks) + + save_mlx_h3_checkpoint(dit, tmp_path) + loaded = load_mlx_h3_checkpoint(tmp_path) + assert loaded.adaln_rank == rank + assert len(loaded.blocks) == 42 + np.testing.assert_array_equal(loaded._adaln_cache.timesteps, timesteps) + np.testing.assert_array_equal(np.array(loaded._adaln_cache.norm_out_scale), np.array(cache.norm_out_scale)) diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py new file mode 100644 index 0000000000..288b3bf465 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py @@ -0,0 +1,176 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Native MLX NVFP4 encoder storage and residency regression checks.""" +import json +from types import SimpleNamespace + +import numpy as np +import pytest + +mx = pytest.importorskip("mlx.core") +from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + NVFP4Matrix, ResidentNVFP4MiniMaxH3TextConditioner, _ResidentNVFP4Index, + _ShardIndex, export_mlx_h3_nvfp4_encoder, unswizzle_nvfp4_scales, +) +from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline + + +def _swizzle(values): + # Independent coordinate mapping of FlashInfer's 128-row, four-column tiles. + rows, cols = values.shape + padded = np.zeros((-(-rows // 128) * 128, -(-cols // 4) * 4), np.uint8) + padded[:rows, :cols] = values + output = np.empty(padded.size, np.uint8) + for r in range(padded.shape[0]): + for c in range(padded.shape[1]): + address = ((((r // 128) * (padded.shape[1] // 4) + c // 4) * 32 + + r % 32) * 4 + (r % 128) // 32) * 4 + c % 4 + output[address] = padded[r, c] + return output + + +def _decode_e4m3(values): + sign = np.where(values & 128, -1.0, 1.0) + exponent = (values >> 3) & 15 + fraction = values & 7 + return sign * np.where(exponent == 0, fraction * 2.0**-9, + (1.0 + fraction / 8.0) * 2.0**(exponent.astype(int) - 7)) + + +def test_padded_scale_layout_round_trip(): + rng = np.random.default_rng(11) + values = rng.integers(0, 127, (140, 7), dtype=np.uint8) + np.testing.assert_array_equal(unswizzle_nvfp4_scales(_swizzle(values), 140, 7), values) + with pytest.raises(ValueError, match="bytes"): + unswizzle_nvfp4_scales(np.zeros(1, np.uint8), 140, 7) + + +@pytest.mark.parametrize("global_scale", [0.5, 4.0]) +@pytest.mark.parametrize("native_cache", [False, True]) +def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, global_scale, native_cache): + from safetensors.numpy import save_file + + rng = np.random.default_rng(3) + packed = rng.integers(0, 256, (128, 64), dtype=np.uint8) + scales = rng.integers(24, 96, (128, 8), dtype=np.uint8) + prefix = "model.language_model.layers.0.self_attn.q_proj" + save_file({prefix + ".weight_packed": packed, + prefix + ".weight_scale": _swizzle(scales), + prefix + ".weight_global_scale": np.array([global_scale], np.float32)}, + tmp_path / "model.safetensors") + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + if native_cache: + _write_encoder_config(tmp_path) + cache_dir = export_mlx_h3_nvfp4_encoder(tmp_path, tmp_path / "cache") + cached = _ResidentNVFP4Index.from_mlx_checkpoint(cache_dir) + for key, original in index.weights.items(): + value = cached.get_mlx(key) + np.testing.assert_array_equal(np.array(value.weight), np.array(original.weight)) + np.testing.assert_array_equal(np.array(value.scales), np.array(original.scales)) + assert value.global_scale == original.global_scale + index.close() + index = cached + weight = index.get_mlx(prefix + ".weight") + assert isinstance(weight, NVFP4Matrix) + assert weight.weight.dtype == mx.uint32 + assert weight.scales.dtype == mx.uint8 + lut = np.array([0, .5, 1, 1.5, 2, 3, 4, 6, 0, -.5, -1, -1.5, -2, -3, -4, -6], np.float32) + dense = np.stack((lut[packed & 15], lut[packed >> 4]), axis=-1).reshape(128, 128) + dense *= np.repeat(_decode_e4m3(scales), 16, axis=1) / global_scale + x = rng.standard_normal((3, 128)).astype(np.float32) + np.testing.assert_allclose(np.array(weight.matmul(mx.array(x))), x @ dense.T, rtol=3e-5, atol=1e-3) + index.close() + assert not index.weights + + +def _write_encoder_config(path): + (path / "config.json").write_text(json.dumps({"quantization_config": { + "quant_method": "nvfp4", "fmt": "e2m1", "group_size": 16, + "scale_fmt": "e4m3", "scale_layout": "128x4", "activation_scheme": "dynamic", + }})) + + +@pytest.mark.parametrize("native_cache", [False, True]) +def test_resident_embedding_keeps_bf16_storage(tmp_path, native_cache): + torch = pytest.importorskip("torch") + from safetensors.torch import save_file + + key = "model.language_model.embed_tokens.weight" + table = torch.arange(60).reshape(10, 6).to(torch.bfloat16) + save_file({key: table}, tmp_path / "model.safetensors") + if native_cache: + _write_encoder_config(tmp_path) + cache_dir = export_mlx_h3_nvfp4_encoder(tmp_path, tmp_path / "cache") + index = _ResidentNVFP4Index.from_mlx_checkpoint(cache_dir) + with pytest.raises(FileExistsError, match="empty"): + export_mlx_h3_nvfp4_encoder(tmp_path, cache_dir) + else: + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + assert index.get_mlx(key).dtype == mx.bfloat16 + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + conditioner.index = index + np.testing.assert_array_equal(np.array(conditioner._embed_tokens([7, 1])), table[[7, 1]].float().numpy()) + conditioner.close() + + +def _pipeline(): + pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) + pipeline.resident = True + pipeline._resident_components = {} + pipeline.dit_checkpoint = "tiny" + pipeline.model_root = __import__("pathlib").Path("tiny") + pipeline.vae_dtype = "fp16" + return pipeline + + +def test_resident_preload_reuses_models_and_evaluates_audio(monkeypatch): + import fastvideo.mlx_runtime.minimax_h3_pipeline as module + import fastvideo.mlx_runtime.minimax_h3_audio_vae as audio + import fastvideo.mlx_runtime.minimax_h3_video_vae as video + + pipeline = _pipeline() + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + conditioner.index = SimpleNamespace(close=lambda: None) + monkeypatch.setattr(pipeline, "_load_conditioner", lambda: conditioner) + calls = [] + def load_dit(path): + calls.append(path) + return SimpleNamespace(weights={"x": mx.ones((1,))}, blocks=[], refiner=[], _adaln_cache=None) + monkeypatch.setattr(module, "load_mlx_h3_checkpoint", load_dit) + monkeypatch.setattr(video, "mlx_h3_video_vae_from_dir", lambda *a, **k: object()) + decoder = SimpleNamespace(weights={"x": mx.ones((2,))}) + monkeypatch.setattr(audio, "mlx_h3_audio_vae_from_dir", lambda *a, **k: decoder) + pipeline.prepare_resident() + pipeline.prepare_resident() + assert calls == ["tiny"] + assert set(pipeline._resident_components) == {"conditioner", "dit", "video_vae", "audio_vae"} + pipeline.close() + assert not pipeline._resident_components + + +def test_failed_preload_releases_encoder(monkeypatch): + import fastvideo.mlx_runtime.minimax_h3_pipeline as module + + pipeline = _pipeline() + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + closed = [] + conditioner.index = SimpleNamespace(close=lambda: closed.append(True)) + monkeypatch.setattr(pipeline, "_load_conditioner", lambda: conditioner) + def fail(path): + raise RuntimeError("out of memory") + monkeypatch.setattr(module, "load_mlx_h3_checkpoint", fail) + with pytest.raises(RuntimeError, match="out of memory"): + pipeline.prepare_resident() + assert closed == [True] + assert not pipeline._resident_components + + +def test_single_shard_omits_unused_language_layers_and_vision(tmp_path): + from safetensors.numpy import save_file + + kept = "model.language_model.layers.49.input_layernorm.weight" + dropped = "model.language_model.layers.50.input_layernorm.weight" + save_file({kept: np.ones(8, np.float32), dropped: np.ones(8, np.float32), + "model.visual.weight": np.ones((8, 8), np.float32)}, tmp_path / "model.safetensors") + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + assert set(index.weights) == {kept} + index.close() diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py new file mode 100644 index 0000000000..376478624a --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py @@ -0,0 +1,24 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Changing the gather memory budget must preserve selected-tile attention.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +import fastvideo.mlx_runtime.minimax_h3_vsa as vsa + + +def test_gather_budget_preserves_attention_with_partial_tiles(monkeypatch): + geom = vsa.build_h3_tile_geometry((7, 67), (6, 8, 8), 64) + rng = np.random.default_rng(2026) + shape = (geom.padded_length, 2, 128) + q, k, value = [mx.array(rng.normal(size=shape).astype(np.float32)).astype(mx.bfloat16) for _ in range(3)] + # Reverse video order also exercises the selected-key order, not a dense mask. + selected = np.array([0, 1, geom.num_tiles - 1, geom.num_prefix_tiles], dtype=np.int32) + idx = mx.array(np.broadcast_to(selected, (2, geom.num_video_tiles, selected.size)).copy()) + monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 2 * 1024**3) + expected = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) + mx.eval(expected) + monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 1) + actual = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) + mx.eval(actual) + np.testing.assert_array_equal(np.array(actual.astype(mx.float32)), np.array(expected.astype(mx.float32))) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index b63b50a067..e4d08f830f 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -5,6 +5,8 @@ # Copyright 2025 The FastVideo Authors. from __future__ import annotations + +import os import contextlib import re from collections.abc import Callable, Generator diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index 21be0864a9..e4f7a08bf3 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -16,6 +16,7 @@ import torch.nn.functional as F from torch.utils.checkpoint import checkpoint +import fastvideo.envs as envs from fastvideo.attention import get_attn_backend from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig from fastvideo.platforms import AttentionBackendEnum @@ -525,7 +526,7 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: def _tile_batch_size() -> int: """Spatial tiles decoded per decoder call (``FASTVIDEO_H3_VAE_TILE_BATCH``, default 1 = per tile).""" - return max(1, int(os.environ.get("FASTVIDEO_H3_VAE_TILE_BATCH", "1"))) + return max(1, envs.FASTVIDEO_H3_VAE_TILE_BATCH.get()) def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool: diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 7dbb74d1c1..0bd2641e5f 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -466,6 +466,9 @@ def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: def _move_module(self, module: Any, device: str | torch.device) -> bool: if _module_has_dtensor_params(module): return False + if not callable(getattr(module, "named_parameters", None)): + module.to(device) + return True if os.environ.get("FASTVIDEO_H3_PINNED_SWAP", "1") == "1": _pinned_swap(module, torch.device(device)) else: diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index c87d946f65..71d9dd0133 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -191,8 +191,14 @@ def _execute( if torch.cuda.is_available(): gib = 1024**3 logger.info("[%s] Memory peak_allocated=%.2f GiB reserved=%.2f GiB resident_after=%.2f GiB", - stage_name, torch.cuda.max_memory_allocated() / gib, - torch.cuda.memory_reserved() / gib, torch.cuda.memory_allocated() / gib) + stage_name, + torch.cuda.max_memory_allocated() / gib, + torch.cuda.memory_reserved() / gib, + torch.cuda.memory_allocated() / gib) + batch.logging_info.add_stage_metric(stage_key, "peak_allocated_mb", + torch.cuda.max_memory_allocated() / 1024**2) + batch.logging_info.add_stage_metric(stage_key, "peak_reserved_mb", + torch.cuda.max_memory_reserved() / 1024**2) torch.cuda.reset_peak_memory_stats() batch.logging_info.add_stage_execution_time(stage_key, execution_time) batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py index e009505547..6e59682545 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py @@ -21,6 +21,7 @@ MiniMaxH3MLXPipeline, _adaln_schedule_union, _center_crop_frames, + _configure_metal_memory_limits, _default_metal_wired_limit_gib, _preflight_media_dependencies, _validate_checkpoint_step_ladder, @@ -143,3 +144,26 @@ def fail(*_args, **_kwargs): assert not output.with_suffix(".tmp.mp4").exists() assert not output.with_suffix(".tmp.wav").exists() + + +def test_explicit_wired_limit_uses_wired_api_separately_from_allocator(): + calls = [] + fake = SimpleNamespace( + metal=SimpleNamespace(device_info=lambda: {"memory_size": 36 * 2**30}), + set_memory_limit=lambda size: calls.append(("allocator", size)), + set_wired_limit=lambda size: calls.append(("wired", size)) or 0, + ) + _configure_metal_memory_limits(fake, 27.0) + assert calls == [("allocator", 30 * 2**30), ("wired", 27 * 2**30)] + + +def test_explicit_wired_limit_failure_is_not_silently_ignored(): + def reject(size): + raise ValueError("exceeds system wired limit") + fake = SimpleNamespace(set_wired_limit=reject) + with pytest.raises(ValueError, match="system wired limit"): + _configure_metal_memory_limits(fake, 31.0) + with pytest.raises(ValueError, match="finite and positive"): + _configure_metal_memory_limits(fake, float("nan")) + with pytest.raises(RuntimeError, match="cannot set"): + _configure_metal_memory_limits(SimpleNamespace(), 27.0) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index ca3fe3d11f..f89cc584bd 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -328,3 +328,23 @@ def test_converter_continues_past_mismatched_existing_format(tmp_path, monkeypat converter.main() assert saved == ["int6"] assert (existing / h3.H3_WEIGHTS_FILENAME).read_bytes() == b"existing" + + +@pytest.mark.parametrize("dtype", [mx.float16, mx.bfloat16]) +@pytest.mark.parametrize("exempt", [False, True]) +def test_simd_partial_key_chunks_match_reference(dtype, exempt): + """Nonuniform scores and partial tiles exercise all four 8-key fragments.""" + _require_metal() + geometry = vsa.build_h3_tile_geometry((7, 5), (5, 3, 7), 64) + mx.random.seed(3026) + q, k, value = [mx.random.normal((geometry.total_seq_length, 2, 128)).astype(dtype) + for _ in range(3)] + expected = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="reference") + stats = vsa.MiniMaxH3VSAStats() + actual = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="simd", stats=stats) + mx.eval(actual, expected) + assert stats.impl == "simd" and stats.dense_fallback_reason is None + assert mx.all(mx.isfinite(actual)).item() + # The Metal kernel uses FP32 accumulation with a different reduction order. + np.testing.assert_allclose(np.asarray(actual.astype(mx.float32)), + np.asarray(expected.astype(mx.float32)), atol=.01, rtol=.01) diff --git a/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py new file mode 100644 index 0000000000..7d34e346f4 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU regression for the ModelOpt activation-scale export contract.""" +import importlib.util +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + + +@pytest.fixture +def converter(monkeypatch): + path = Path(__file__).resolve().parents[4] / 'scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py' + spec = importlib.util.spec_from_file_location('h3_modelopt_converter', path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + packed = torch.zeros((2, 8), dtype=torch.uint8) + scales = torch.ones((2, 1), dtype=torch.float8_e4m3fn) + def quantize(*args, **kwargs): + return packed, scales + monkeypatch.setattr(module, '_flashinfer', lambda: ( + SimpleNamespace(layout_128x4=0), None, quantize, lambda value: value)) + return module + + +def convert(converter, input_scale): + return converter.convert_modelopt_linear( + torch.zeros((2, 8), dtype=torch.uint8), + torch.ones((2, 1), dtype=torch.float8_e4m3fn), + torch.tensor(0.25), 'cpu', input_scale=input_scale)[0] + + +def test_modelopt_activation_scale_is_preserved_as_reciprocal(converter): + buffers = convert(converter, torch.tensor(2.0)) + assert buffers['_nvfp4_input_global_sf'].dtype == torch.float32 + assert buffers['_nvfp4_input_global_sf'].shape == () + assert buffers['_nvfp4_input_global_sf'].item() == 0.5 + assert buffers['_nvfp4_alpha'].item() == 0.25 + + +def test_missing_modelopt_activation_scale_keeps_legacy_export(converter): + assert '_nvfp4_input_global_sf' not in convert(converter, None) + + +@pytest.mark.parametrize('value', [0.0, -1.0, float('nan'), float('inf')]) +def test_invalid_modelopt_activation_scale_is_rejected(converter, value): + with pytest.raises(ValueError, match='finite positive scalar'): + convert(converter, torch.tensor(value)) diff --git a/fastvideo/tests/worker/test_ray_distributed_executor.py b/fastvideo/tests/worker/test_ray_distributed_executor.py index 0bf8b4d1e3..ffb3e48421 100644 --- a/fastvideo/tests/worker/test_ray_distributed_executor.py +++ b/fastvideo/tests/worker/test_ray_distributed_executor.py @@ -36,3 +36,17 @@ def test_ray_log_queue_stays_on_the_driver() -> None: executor.clear_log_queue() assert executor._log_queue is None assert "log_queue" in signature(Executor.set_log_queue).parameters + + +def test_ray_carries_h3_performance_switches() -> None: + import fastvideo.envs as envs + from fastvideo.worker.ray_env import get_env_vars_to_copy + + with (envs.FASTVIDEO_H3_VAE_TILE_BATCH.override(8), + envs.FASTVIDEO_NVFP4_MM_BACKEND.override("cutlass"), + envs.FASTVIDEO_VSA_TRITON.override(True)): + copied = get_env_vars_to_copy() + assert {"FASTVIDEO_H3_VAE_TILE_BATCH", "FASTVIDEO_NVFP4_MM_BACKEND", "FASTVIDEO_VSA_TRITON"} <= copied + assert envs.FASTVIDEO_H3_VAE_TILE_BATCH.get() == 8 + assert envs.FASTVIDEO_NVFP4_MM_BACKEND.get() == "cutlass" + assert envs.FASTVIDEO_VSA_TRITON.get() is True diff --git a/fastvideo/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index 4fdcae2c5e..beb2af9385 100644 --- a/fastvideo/worker/gpu_worker.py +++ b/fastvideo/worker/gpu_worker.py @@ -1,4 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 +import os from typing import Any, cast import torch @@ -21,7 +22,6 @@ def _log_cuda_device_uuid(rank: int, device: torch.device) -> None: logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False) - def _log_pipeline_memory(pipeline) -> None: """Debug (FASTVIDEO_MEMORY_REPORT=1): bytes held per pipeline component, by device and dtype, plus the largest tensors, so the resident footprint can be attributed before choosing offload placements.""" @@ -42,14 +42,18 @@ def _log_pipeline_memory(pipeline) -> None: largest.append((nbytes, tname, key)) largest.sort(reverse=True) total = sum(by_kind.values()) - logger.info("MEMREPORT %s total=%.2f GiB %s", name, total / gib, - {k: round(v / gib, 2) for k, v in sorted(by_kind.items(), key=lambda kv: -kv[1])}) + logger.info("MEMREPORT %s total=%.2f GiB %s", name, total / gib, { + k: round(v / gib, 2) + for k, v in sorted(by_kind.items(), key=lambda kv: -kv[1]) + }) for nbytes, tname, key in largest[:8]: logger.info("MEMREPORT %s %.3f GiB %s %s", name, nbytes / gib, key, tname) if torch.cuda.is_available(): - logger.info("MEMREPORT cuda allocated=%.2f GiB reserved=%.2f GiB", torch.cuda.memory_allocated() / gib, + logger.info("MEMREPORT cuda allocated=%.2f GiB reserved=%.2f GiB", + torch.cuda.memory_allocated() / gib, torch.cuda.memory_reserved() / gib) + class Worker: def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str): diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index db96f413c2..4a43adaea2 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -18,6 +18,17 @@ INT8/INT6/INT4 grid, and record ``vsa.capable`` in the manifest. Write VSA checkpoints to a new directory — do not overwrite an existing dense export. +An FP8 transformer with per-channel ``weight_scale`` can also be a source. +The loader dequantizes each FP8 matrix before applying the requested MLX +quantization. This saves download bytes but quantizes twice; compare its clips +with the BF16-sourced export before using it for release. + +``--formats "mxfp8 mxfp4 nvfp4"`` tries native MLX floating-point quantized +storage and matrix multiplication. These formats are experimental and require +operator support from the installed MLX build. This converts BF16 or FP8 source +weights; it does not import CUDA-packed NVFP4 DiT exports. The default formats +remain affine INT8/INT6/INT4. + python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ --out ~/models/FastH3-MLX-vsa \\ @@ -53,8 +64,8 @@ logger = init_logger(__name__) -SUPPORTED_FORMATS = ("int8", "int6", "int4") -DEFAULT_FORMATS = " ".join(SUPPORTED_FORMATS) +SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8", "mxfp4", "nvfp4") +DEFAULT_FORMATS = "int8 int6 int4" def _adaln_cache_timesteps(model_root: str | Path | None = None) -> np.ndarray: @@ -90,6 +101,10 @@ def parse_args() -> argparse.Namespace: help=("retain and quantize transformer_blocks.*.attn.to_gate_compress.weight " "(required for MLX VSA inference; omitted by dense conversion)"), ) + parser.add_argument("--nvfp4-conditioner-root", type=Path, + help="also cache the released packed encoder in native MLX layout, without requantization") + parser.add_argument("--nvfp4-conditioner-out", type=Path, + help="empty encoder-cache output directory; defaults to OUT/nvfp4-encoder") return parser.parse_args() @@ -144,6 +159,15 @@ def main() -> None: if hasattr(mx, "clear_cache"): mx.clear_cache() + if args.nvfp4_conditioner_root is not None: + from fastvideo.mlx_runtime.minimax_h3_conditioner import export_mlx_h3_nvfp4_encoder + + encoder_out = args.nvfp4_conditioner_out or out_base / "nvfp4-encoder" + started = time.perf_counter() + export_mlx_h3_nvfp4_encoder(args.nvfp4_conditioner_root, encoder_out) + print(f"[encoder] cached packed NVFP4 encoder in {time.perf_counter() - started:.1f}s at {encoder_out}", + flush=True) + if __name__ == "__main__": main() diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py index 8b2ff1b448..4f0b3728b7 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -13,10 +13,11 @@ ``load_minimax_h3_nvfp4_dit_export``) stores `` ::_nvfp4_weight`` (same bytes), ``::_nvfp4_weight_scale`` (the same E4M3 bytes in FlashInfer's 128x4 swizzled layout), ``::_nvfp4_alpha`` (= ``weight_scale_2``) and -``::_weight_global_sf`` (= 1 / ``weight_scale_2``). The calibrated weights are -therefore carried over bit for bit; only the activation scale changes, because -FastVideo quantizes activations per call with a unit global scale and the -static ``input_scale`` is dropped. +``::_weight_global_sf`` (= 1 / ``weight_scale_2``), and +``::_nvfp4_input_global_sf`` (= 1 / ``input_scale``). The calibrated weight +bytes and activation scale are preserved. Dropping the activation scale would +replace its calibrated range with a unit global scale, clipping inputs above +2688. ``--quantize-attention`` additionally quantizes the dense BF16 attention projections (``attn.to_{q,k,v,out}``) of every main block exactly as @@ -105,7 +106,8 @@ def probe(buffers: dict[str, torch.Tensor], reference: torch.Tensor, rows: int = return ((out.float() - ref).norm() / ref.norm()).item() -def convert_modelopt_linear(weight, scale, scale_2, device) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: +def convert_modelopt_linear(weight, scale, scale_2, device, input_scale=None + ) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: """Carry the calibrated bytes over; return (buffers, dequantized weight, scale-byte agreement). The agreement compares the swizzled ModelOpt scales with the scales @@ -128,6 +130,11 @@ def convert_modelopt_linear(weight, scale, scale_2, device) -> tuple[dict[str, t "_nvfp4_alpha": scale_2.clone(), "_weight_global_sf": (1.0 / scale_2).to(torch.bfloat16), } + if input_scale is not None: + value = input_scale.to(dtype=torch.float32) + if value.numel() != 1 or not bool(torch.isfinite(value).all()) or not bool((value > 0).all()): + raise ValueError("ModelOpt input_scale must be a finite positive scalar") + buffers["_nvfp4_input_global_sf"] = value.reshape(()).reciprocal().to(device=device) return buffers, reference, agreement @@ -189,7 +196,9 @@ def main() -> None: else: buffers, reference, agreement = convert_modelopt_linear(tensor(f"{prefix}.weight"), tensor(f"{prefix}.weight_scale"), - tensor(f"{prefix}.weight_scale_2"), device) + tensor(f"{prefix}.weight_scale_2"), device, + input_scale=tensor(f"{prefix}.input_scale") + if f"{prefix}.input_scale" in weight_map else None) agreements.append(agreement) error = probe(buffers, reference) worst = max(worst, error) @@ -233,7 +242,7 @@ def main() -> None: shutil.copy2(extra, args.dst / extra.name) print(json.dumps({"exported_linears": len(modelopt) + len(attention), "modelopt_linears": len(modelopt), "quantized_dense_linears": len(attention), "worst_probe_error": round(worst, 4), - "static_activation_scales": len(attention) + len(modelopt) if amax_table else 0, + "static_activation_scales": sum(k.endswith("::_nvfp4_input_global_sf") for k in export), "gate_linears": sum(1 for p in attention if _BLOCK_GATE.match(p)), "min_scale_byte_agreement": round(min(agreements), 4) if agreements else None, "mean_scale_byte_agreement": round(sum(agreements) / len(agreements), 4) if agreements else None,