Skip to content

minimax_h3_vsa

MiniMax H3 Video Sparse Attention for the native MLX runtime.

Ports the packed-sequence VSA-H3 contract from fastvideo/attention/backends/video_sparse_attn_h3.py:

  • Tiles are [segment-pure prefix chunks] + [3D video tiles].
  • Tile sizes 64 (4, 4, 4) and 256 (4, 8, 8).
  • Per-head pooled Q/K scoring and top-k routing.
  • Prefix queries are always dense; prefix keys are exempt (always kept) or compete (FLOP-matched top-k).
  • Optional dense-first steps and per-layer dense overrides.
  • Trained to_gate_compress pooled-compression branch.

Execution backends:

  • reference (auto default) — grouped gather plus batched mx.fast.scaled_dot_product_attention. Correctness baseline.
  • simd — opt-in SIMD-group 8x8 matrix operations over the reference tile map for tile size 64 and head dimension 128. Unsupported shapes and kernel failures fall back to the reference backend.

Dense fused SDPA remains the default when VSA is disabled, the geometry is unsupported, or a dense-only checkpoint is loaded.

Classes

fastvideo.mlx_runtime.minimax_h3_vsa.DenseOnlyVSACheckpointError

Bases: ValueError

VSA was requested for a checkpoint that dropped the gate weights.

fastvideo.mlx_runtime.minimax_h3_vsa.MiniMaxH3VSAConfig dataclass

MiniMaxH3VSAConfig(enabled: bool = False, sparsity: float = 0.9, tile_size: int = 64, prefix_mode: PrefixMode = 'exempt', dense_first_n_steps: int = 0, dense_layers: tuple[int, ...] = (), impl: VSAImpl = 'auto')

Runtime VSA knobs. Defaults preserve dense MLX H3 behavior.

fastvideo.mlx_runtime.minimax_h3_vsa.MiniMaxH3VSAGeometry dataclass

MiniMaxH3VSAGeometry(prefix_segments: tuple[int, ...], dit_seq_shape: tuple[int, int, int], tile_shape: tuple[int, int, int], tile_elems: int, total_seq_length: int, num_prefix_tiles: int, num_video_tiles: int, variable_block_sizes: ndarray, untile_combined_index: ndarray, tile_partition_indices: ndarray)

Packed-sequence tile map shared by routing, reference, and Metal paths.

fastvideo.mlx_runtime.minimax_h3_vsa.MiniMaxH3VSAStats dataclass

MiniMaxH3VSAStats(configured_sparsity: float = 0.0, layer_sparsity: float = 0.0, tile_size: int = 64, prefix_mode: str = 'exempt', impl: str = 'dense', num_prefix_tiles: int = 0, num_video_tiles: int = 0, video_keep: float = 0.0, achieved_sparsity: float = 0.0, dense_fallback_reason: str | None = None, attention_calls: int = 0, sparse_calls: int = 0, impl_counts: dict[str, int] = dict(), fallback_reasons: list[str] = list())

Filled during a sparse forward so the pipeline can report achieved sparsity.

Methods:

fastvideo.mlx_runtime.minimax_h3_vsa.MiniMaxH3VSAStats.record
record(call: MiniMaxH3VSAStats) -> None

Aggregate equally sized video-query tile maps across blocks and steps.

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def record(self, call: MiniMaxH3VSAStats) -> None:
    """Aggregate equally sized video-query tile maps across blocks and steps."""
    self.attention_calls += 1
    self.sparse_calls += int(call.layer_sparsity > 0.0)
    self.impl_counts[call.impl] = self.impl_counts.get(call.impl, 0) + 1
    self.impl = next(iter(self.impl_counts)) if len(self.impl_counts) == 1 else "mixed"
    for name in ("layer_sparsity", "video_keep", "achieved_sparsity"):
        old = getattr(self, name)
        setattr(self, name, old + (getattr(call, name) - old) / self.attention_calls)
    self.num_prefix_tiles = call.num_prefix_tiles
    self.num_video_tiles = call.num_video_tiles
    if call.dense_fallback_reason and call.dense_fallback_reason not in self.fallback_reasons:
        self.fallback_reasons.append(call.dense_fallback_reason)
    self.dense_fallback_reason = "; ".join(self.fallback_reasons) or None

Functions:

fastvideo.mlx_runtime.minimax_h3_vsa.build_block_mask

build_block_mask(scores: ndarray, num_prefix_tiles: int, num_video_tiles: int, sparsity: float, exempt: bool) -> ndarray

scores: [..., n_tiles, n_tiles] -> bool mask, same shape.

Mirrors _build_block_mask in the PyTorch H3 backend.

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def build_block_mask(
    scores: np.ndarray,
    num_prefix_tiles: int,
    num_video_tiles: int,
    sparsity: float,
    exempt: bool,
) -> np.ndarray:
    """scores: [..., n_tiles, n_tiles] -> bool mask, same shape.

    Mirrors ``_build_block_mask`` in the PyTorch H3 backend.
    """
    n_tiles = scores.shape[-1]
    k_vid = compute_topk(sparsity, num_video_tiles)
    if k_vid == num_video_tiles:
        return np.ones_like(scores, dtype=bool)
    mask = np.zeros_like(scores, dtype=bool)
    if exempt or num_prefix_tiles == 0:
        video_cols = scores[..., num_prefix_tiles:]
        idx = np.argsort(-video_cols, axis=-1)[..., :k_vid] + num_prefix_tiles
        np.put_along_axis(mask, idx, True, axis=-1)
        mask[..., :num_prefix_tiles] = True
    else:
        k_total = min(k_vid + num_prefix_tiles, n_tiles)
        idx = np.argsort(-scores, axis=-1)[..., :k_total]
        np.put_along_axis(mask, idx, True, axis=-1)
    mask[..., :num_prefix_tiles, :] = True
    return mask

fastvideo.mlx_runtime.minimax_h3_vsa.build_h3_tile_geometry

build_h3_tile_geometry(prefix_segments: tuple[int, ...], dit_seq_shape: tuple[int, int, int], tile_size: int = 64) -> MiniMaxH3VSAGeometry

Tile the packed sequence: segment-pure prefix chunks, then video tiles.

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def build_h3_tile_geometry(
    prefix_segments: tuple[int, ...],
    dit_seq_shape: tuple[int, int, int],
    tile_size: int = 64,
) -> MiniMaxH3VSAGeometry:
    """Tile the packed sequence: segment-pure prefix chunks, then video tiles."""
    tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
    if tile_shape is None:
        raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
    if len(dit_seq_shape) != 3 or any(axis <= 0 for axis in dit_seq_shape):
        raise ValueError(f"VSA-H3 video axes must be positive, got {dit_seq_shape}.")
    tile_elems = math.prod(tile_shape)
    prefix_segments = tuple(int(segment) for segment in prefix_segments if int(segment) > 0)
    prefix_len = sum(prefix_segments)

    prefix_sizes: list[int] = []
    for segment in prefix_segments:
        full, rem = divmod(segment, tile_elems)
        prefix_sizes.extend([tile_elems] * full)
        if rem:
            prefix_sizes.append(rem)
    num_prefix_tiles = len(prefix_sizes)

    video_sizes = _video_tile_sizes(dit_seq_shape, tile_shape)
    num_video_tiles = int(video_sizes.size)
    video_indices = _video_tile_partition_indices(dit_seq_shape, tile_shape) + prefix_len
    tile_partition_indices = np.concatenate(
        [np.arange(prefix_len, dtype=np.int64), video_indices],
        axis=0,
    )
    variable_block_sizes = np.concatenate(
        [np.asarray(prefix_sizes, dtype=np.int64),
         video_sizes.astype(np.int64)],
        axis=0,
    )
    non_pad_index = _non_pad_index(variable_block_sizes, tile_elems)
    untile_combined_index = non_pad_index[np.argsort(tile_partition_indices, kind="stable")]
    validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
    return MiniMaxH3VSAGeometry(
        prefix_segments=prefix_segments,
        dit_seq_shape=dit_seq_shape,
        tile_shape=tile_shape,
        tile_elems=tile_elems,
        total_seq_length=prefix_len + math.prod(dit_seq_shape),
        num_prefix_tiles=num_prefix_tiles,
        num_video_tiles=num_video_tiles,
        variable_block_sizes=variable_block_sizes,
        untile_combined_index=untile_combined_index,
        tile_partition_indices=tile_partition_indices,
    )

fastvideo.mlx_runtime.minimax_h3_vsa.compute_topk

compute_topk(sparsity: float, num_blocks: int) -> int

Blocks to keep for a sparsity level, clamped to [1, num_blocks].

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def compute_topk(sparsity: float, num_blocks: int) -> int:
    """Blocks to keep for a sparsity level, clamped to [1, num_blocks]."""
    if num_blocks <= 0:
        return 0
    return max(1, min(math.ceil((1.0 - sparsity) * num_blocks), num_blocks))

fastvideo.mlx_runtime.minimax_h3_vsa.geometry_is_supported

geometry_is_supported(prefix_segments: tuple[int, ...], dit_seq_shape: tuple[int, int, int], tile_size: int) -> str | None

Return a fallback reason, or None when VSA can run.

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def geometry_is_supported(prefix_segments: tuple[int, ...], dit_seq_shape: tuple[int, int, int],
                          tile_size: int) -> str | None:
    """Return a fallback reason, or None when VSA can run."""
    try:
        build_h3_tile_geometry(prefix_segments, dit_seq_shape, tile_size)
    except ValueError as error:
        return str(error)
    return None

fastvideo.mlx_runtime.minimax_h3_vsa.h3_vsa_attention

h3_vsa_attention(query, key, value, geometry: MiniMaxH3VSAGeometry, *, sparsity: float, exempt: bool = True, gate_compress=None, impl: VSAImpl = 'auto', stats: MiniMaxH3VSAStats | None = None)

Packed [S, H, D] VSA attention. Falls back to dense SDPA when sparsity is 0.

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def h3_vsa_attention(
    query,
    key,
    value,
    geometry: MiniMaxH3VSAGeometry,
    *,
    sparsity: float,
    exempt: bool = True,
    gate_compress=None,
    impl: VSAImpl = "auto",
    stats: MiniMaxH3VSAStats | None = None,
):
    """Packed ``[S, H, D]`` VSA attention. Falls back to dense SDPA when sparsity is 0."""
    import mlx.core as mx

    _, heads, dim = query.shape
    scale = dim**-0.5
    if stats is not None:
        stats.configured_sparsity = sparsity
        stats.layer_sparsity = sparsity
        stats.tile_size = geometry.tile_elems
        stats.prefix_mode = "exempt" if exempt else "compete"
        stats.num_prefix_tiles = geometry.num_prefix_tiles
        stats.num_video_tiles = geometry.num_video_tiles

    if sparsity <= 0.0 and gate_compress is None:
        if stats is not None:
            stats.impl = "dense"
            stats.achieved_sparsity = 0.0
            stats.video_keep = geometry.num_video_tiles
        return _dense_sdpa(query, key, value, scale)

    q_tiled = _tile_hidden(query, geometry)
    k_tiled = _tile_hidden(key, geometry)
    v_tiled = _tile_hidden(value, geometry)
    q_pooled = _pool_tiles(q_tiled, geometry.variable_block_sizes, geometry.tile_elems)
    k_pooled = _pool_tiles(k_tiled, geometry.variable_block_sizes, geometry.tile_elems)
    scores = (q_pooled @ k_pooled.transpose(0, 2, 1)) / (dim**0.5)

    k_vid = compute_topk(sparsity, geometry.num_video_tiles)
    if stats is not None:
        stats.video_keep = k_vid if sparsity > 0.0 else geometry.num_video_tiles
        stats.achieved_sparsity = (1.0 - (k_vid / geometry.num_video_tiles) if geometry.num_video_tiles else 0.0)

    if sparsity <= 0.0:
        out_tiled = _tile_hidden(_dense_sdpa(query, key, value, scale), geometry)
        chosen = "dense"
    else:
        block_idx, block_num = _block_indices_from_scores(
            scores,
            geometry.num_prefix_tiles,
            geometry.num_video_tiles,
            sparsity,
            exempt,
        )
        if stats is not None and not exempt:
            stats.video_keep = float(mx.mean(mx.sum(block_idx >= geometry.num_prefix_tiles, axis=-1)).item())
            stats.achieved_sparsity = 1.0 - stats.video_keep / geometry.num_video_tiles
        chosen = resolve_impl(impl, dim, geometry.tile_elems)
        if stats is not None and impl == "simd" and chosen != "simd":
            from fastvideo.mlx_runtime.minimax_h3_vsa_simd import simd_kernel_error

            stats.dense_fallback_reason = ("SIMD requires tile 64 and head dim 128"
                                           if geometry.tile_elems != 64 or dim != 128 else simd_kernel_error())
        prefix_out = (_dense_sdpa(query[:geometry.prefix_length], key, value, scale)
                      if geometry.prefix_length else query[:0])
        n_prefix_pad = geometry.num_prefix_tiles * geometry.tile_elems
        if chosen == "simd":
            from fastvideo.mlx_runtime.minimax_h3_vsa_simd import disable_simd_kernel, simd_block_sparse

            try:
                video_tiled = simd_block_sparse(q_tiled, k_tiled, v_tiled, block_idx, block_num, geometry, scale)
                # MLX compiles and executes custom Metal kernels lazily.
                mx.eval(video_tiled)
                video_tiled = video_tiled[n_prefix_pad:]
            except Exception as error:  # noqa: BLE001 - keep generation alive on kernel failure
                disable_simd_kernel(error)
                logger.warning_once(f"SIMD VSA kernel failed ({error}); falling back to reference gather+SDPA")
                chosen = "reference"
                if stats is not None:
                    stats.dense_fallback_reason = f"simd kernel failed: {error}"
                video_tiled = _reference_gather_sdpa(q_tiled, k_tiled, v_tiled, block_idx, geometry, scale)
        elif _reference_full_mask_fits(geometry, heads):
            scores_np = np.array(scores, dtype=np.float32)
            mask = build_block_mask(
                scores_np,
                geometry.num_prefix_tiles,
                geometry.num_video_tiles,
                sparsity,
                exempt,
            )
            video_tiled = _reference_token_sdpa(q_tiled, k_tiled, v_tiled, mask, geometry, scale)[n_prefix_pad:]
        else:
            video_tiled = _reference_gather_sdpa(q_tiled, k_tiled, v_tiled, block_idx, geometry, scale)
        video_tiled = video_tiled.astype(query.dtype)
        prefix_tiled = mx.concatenate([
            prefix_out,
            mx.zeros((1, heads, dim), dtype=query.dtype),
        ])[geometry.prefix_gather_index]
        # Prefix query tiles are dense; keep fused-SDPA prefix rows and sparse video tiles.
        out_tiled = mx.concatenate([prefix_tiled, video_tiled], axis=0)

    if gate_compress is not None:
        gate_tiled = _tile_hidden(gate_compress, geometry)
        out_tiled = out_tiled + _gate_compress_output(scores, v_tiled, gate_tiled, geometry).astype(out_tiled.dtype)

    if stats is not None:
        stats.impl = chosen if sparsity > 0.0 else "dense"
    return _untile_hidden(out_tiled, geometry)

fastvideo.mlx_runtime.minimax_h3_vsa.prefix_segments_from_layout

prefix_segments_from_layout(layout: Any, patch_size: tuple[int, int, int]) -> tuple[int, ...]

Segment sizes preceding the generated-video tail, matching the PyTorch stage.

Source code in fastvideo/mlx_runtime/minimax_h3_vsa.py
def prefix_segments_from_layout(layout: Any, patch_size: tuple[int, int, int]) -> tuple[int, ...]:
    """Segment sizes preceding the generated-video tail, matching the PyTorch stage."""
    n_text = int(layout.text_indices.shape[0])
    n_cond = int(layout.num_condition_video_rows)
    n_audio = int(layout.audio_indices.shape[0])
    n_video = ((layout.num_video_latent_frames // patch_size[0]) * (layout.latent_height // patch_size[1]) *
               (layout.latent_width // patch_size[2]))
    if n_text + n_cond + n_audio + n_video != int(layout.sequence_length):
        raise ValueError("VSA-H3 supports the standard [text|cond|audio|video] packing only; "
                         f"segments ({n_text}, {n_cond}, {n_audio}) + video {n_video} do not sum to "
                         f"sequence length {layout.sequence_length}.")
    return n_text, n_cond, n_audio