Skip to content

video_sparse_attn_h3

VSA for MiniMax H3's packed mixed-modality self-attention.

H3 runs one joint bidirectional attention over [text | condition keyframes | audio | generated video], so this backend differs from the Wan-tuned video_sparse_attn:

  • Tiles are [segment-pure prefix chunks] + [3D video tiles]; prefix tiles never straddle segment boundaries. The tile size is selectable at metadata build time: 256 tokens (4,8,8) (default), 128 tokens (4,4,8), or 64 tokens (4,4,4) (see VSA_H3_TILE_SHAPES).
  • The builder takes the packed sequence as ordered segments: dense row counts (text, audio, image references) and one 3-D sparse region per video, wherever it sits. The generated video is the last region; a Ref2VA PDD run adds one earlier region per reference video. Rows are permuted to [dense chunks][region 0 tiles][region 1 tiles]... and untile_combined_index inverts the permutation, so the kernels see one dense prefix followed by video tiles. Every video query keeps its own top-k of EACH region (video_tile_spans / span_sparsities); the reference-video regions may use a different keep rate than the generated video.
  • Selection is pure Python on pooled tile scores; the block-sparse kernel consumes an explicit bool mask, so no kernel changes are needed.
  • The compression branch is gated by to_gate_compress, which the base H3 checkpoint does not carry: the loader zero-initializes it, so untrained inference is exactly pure sparse and finetuning can learn the gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship trained gates, which load and activate the branch.
  • Non-video queries are always dense. Non-video keys are either always-selected for every query ("exempt", default) or compete in top-k under a FLOP-matched budget ("compete") — the ablation axis, switched per request via generate_video(..., vsa_mode=...) (default: exempt). Per-request scheduling knobs (vsa_dense_first_n_steps, vsa_dense_layers) let mixed schedules run the diffuse steps/layers dense while pushing the rest harder.

At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path (FASTVIDEO_VSA_CUTEDSL=1); the Triton 256→64 expansion is the fallback and keeps identical mask semantics. At tile 64 the block map is already at the kernels' native 64-token granularity, so both forward and backward run the Triton block-sparse kernels directly (no expansion, FASTVIDEO_VSA_CUTEDSL does not apply). A third, opt-in route exists for the tile-64 FORWARD only: FASTVIDEO_VSA_SM100A=1 sends no-grad forwards through the data-center Blackwell CUDA block-sparse kernel (fastvideo_kernel.block_sparse_attn_sm100a, which reads a separate q2k_num key-tile count per query tile) when the extension is built, the device is sm_100 or sm_103, and the geometry qualifies. The CUDA kernel assigns adjacent pairs of query tiles to CTAs, so an odd logical tile count receives one internal, zero-valid partner tile for the no-grad call only. Score search, the trained mask, gate-compress, and the returned packed sequence remain on the original logical tiles. Grad-tracking forwards and every backward stay on the Triton kernels. If the env is set but a precondition fails, the route logs one warning and falls back.

Tile 128 has exactly one implementation: the same sm_100a/sm_103a CUDA forward, which carries a 128-token block instantiation. It needs no opt-in, pads an odd tile count with the same zero-valid partner tile, runs no-grad forwards only, and fails closed (no Triton fallback) when the extension or device cannot run it.

Classes

fastvideo.attention.backends.video_sparse_attn_h3.MiniMaxH3VSAImpl

MiniMaxH3VSAImpl(num_heads: int, head_size: int, causal: bool, softmax_scale: float, num_kv_heads: int | None = None, prefix: str = '', **extra_impl_args)

Bases: AttentionImpl

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def __init__(
    self,
    num_heads: int,
    head_size: int,
    causal: bool,
    softmax_scale: float,
    num_kv_heads: int | None = None,
    prefix: str = "",
    **extra_impl_args,
) -> None:
    self.prefix = prefix
    self.layer_idx = layer_idx_from_prefix(prefix, default=-1)
    self.head_size = head_size
    self._sm89_kernel = envs.FASTVIDEO_H3_VSA_SM89_KERNEL.get()
    if self._sm89_kernel not in {"original", "bf16", "int8"}:
        raise ValueError("FASTVIDEO_H3_VSA_SM89_KERNEL must be original, bf16, or int8")
    # Generic torch.compile must not specialize the shared VSA forward on
    # the Python ``layer_idx`` value of each of H3's 50 blocks. This
    # tensor is prepared after weights load and drives only the compiled
    # dense-layer decision; it does not opt the module into sm_100a.
    self._compile_layer_idx: torch.Tensor | None = None
    # None means the regional-compile preparation hook has not run.  The
    # eager path deliberately ignores this cache and preserves its
    # request-time env/probe/fallback behavior; only Dynamo capture reads
    # the prepared, static route.
    self._regional_compile_sm100a_enabled: bool | None = None
    # Tile-128 route per (device, dtype, head size): None when the CUDA
    # kernel can run it, else why not. The kernel's predicate otherwise
    # depends only on the tile-128 buffer contract (contiguous BHSD,
    # 128-token blocks, an even tile count with the partner tile, integer
    # tile sizes), which every call meets, so it is evaluated once.
    self._tile128_route: dict[tuple[torch.device, torch.dtype, int], str | None] = {}

Methods:

fastvideo.attention.backends.video_sparse_attn_h3.MiniMaxH3VSAImpl.prepare_for_compile
prepare_for_compile(device: device) -> None

Tensorize per-layer state shared by every torch.compile route.

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def prepare_for_compile(self, device: torch.device) -> None:
    """Tensorize per-layer state shared by every torch.compile route."""
    self._compile_layer_idx = torch.tensor(self.layer_idx, device=device, dtype=torch.int64)
fastvideo.attention.backends.video_sparse_attn_h3.MiniMaxH3VSAImpl.prepare_for_regional_compile
prepare_for_regional_compile(device: device) -> str | None

Resolve the inference-only sm_100a route before fullgraph capture.

The ordinary eager route probes the environment, extension, device, and tensor contract at every call so it can warn and fall back. Those Python/device-capability checks are not safe inside a regional fullgraph=True block. Probe one representative tile-64 input on the loaded model's device now, then let forward specialize on the resulting plain bool while Dynamo is compiling.

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def prepare_for_regional_compile(self, device: torch.device) -> str | None:
    """Resolve the inference-only sm_100a route before fullgraph capture.

    The ordinary eager route probes the environment, extension, device,
    and tensor contract at every call so it can warn and fall back.  Those
    Python/device-capability checks are not safe inside a regional
    ``fullgraph=True`` block.  Probe one representative tile-64 input on
    the loaded model's device now, then let ``forward`` specialize on the
    resulting plain bool while Dynamo is compiling.
    """
    if self._compile_layer_idx is None:
        self.prepare_for_compile(device)
    requested = envs.FASTVIDEO_VSA_SM100A.get()
    enabled = False
    reason = None if requested else f"{VSA_SM100A_ENV}=1 is required for compile-safe VSA-H3 attention"
    if requested:
        if _sm100a is None:
            reason = "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
        elif not _sm100a_has_compile_safe_mask_route(_sm100a):
            reason = ("neither a native block_sparse_attn_sm100a_from_mask entry nor the raw sm100a "
                      "kernel plus map_to_index compatibility route is installed")
        else:
            # Two 64-token blocks exercise the exact sm_100a inference
            # specialization while keeping the one-time probe tiny.  The
            # kernel predicate checks extension presence, CUDA capability,
            # dtype/layout, head size, block size, and even block count
            # without reading metadata tensor contents.
            probe_query = torch.empty((1, 1, 128, self.head_size), device=device, dtype=torch.bfloat16)
            probe_block_sizes = torch.full((2, ), 64, device=device, dtype=torch.int32)
            reason = _sm100a_unavailable_reason(
                _sm100a,
                probe_query,
                probe_block_sizes,
                grad_mode=False,
            )
            enabled = reason is None

    self._regional_compile_sm100a_enabled = enabled
    if enabled:
        route = ("native fastvideo-kernel mask entry" if callable(
            getattr(_sm100a, "block_sparse_attn_sm100a_from_mask", None)) else
                 "FastVideo compatibility mask adapter")
        logger.info_once(f"VSA-H3 regional compile mask route: {route}")
    if requested and reason is not None:
        logger.warning_once(f"VSA-H3 regional compile is unavailable and will stay eager: {reason}")
    return reason
fastvideo.attention.backends.video_sparse_attn_h3.MiniMaxH3VSAImpl.tile
tile(x: Tensor, attn_metadata: MiniMaxH3VSAMetadata) -> Tensor

Scatter rows into the padded tile buffer (pad positions stay zero).

Without grad tracking the returned tensor aliases the builder-owned buffer; callers must consume it before the next tile() (both call sites in forward() read it immediately). A grad-tracking forward instead receives a fresh buffer and leaves the holder untouched, so the builder never retains autograd state across steps. Odd no-grad sm100a requests (tile 128 always, tile 64 when opted in) carry one additional all-zero tile internally; metadata and all observable outputs retain the logical geometry.

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def tile(self, x: torch.Tensor, attn_metadata: MiniMaxH3VSAMetadata) -> torch.Tensor:
    """Scatter rows into the padded tile buffer (pad positions stay zero).

    Without grad tracking the returned tensor aliases the builder-owned
    buffer; callers must consume it before the next ``tile()`` (both call
    sites in ``forward()`` read it immediately). A grad-tracking forward
    instead receives a fresh buffer and leaves the holder untouched, so
    the builder never retains autograd state across steps. Odd no-grad
    sm100a requests (tile 128 always, tile 64 when opted in) carry one
    additional all-zero tile internally; metadata and all observable
    outputs retain the logical geometry.
    """
    if x.shape[1] != attn_metadata.total_seq_length:
        raise ValueError(f"VSA-H3 metadata was built for sequence length {attn_metadata.total_seq_length}, "
                         f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
                         "routed to the VSA-H3 backend; exclude it from the supported backends.")
    n_tiles = attn_metadata.variable_block_sizes.numel()
    grad_mode = torch.is_grad_enabled() and x.requires_grad
    compiling = torch.compiler.is_compiling()
    regional_compiling = compiling and self._regional_compile_sm100a_enabled is True
    if regional_compiling:
        sm100a_requested = bool(self._regional_compile_sm100a_enabled)
    elif compiling:
        # Training/generic compile keeps the long-standing Triton route.
        sm100a_requested = False
    else:
        sm100a_requested = envs.FASTVIDEO_VSA_SM100A.get()
    # Tile 128 has no route other than the sm100a CUDA kernel.
    sm100a_route = attn_metadata.tile_elems == 128 or (attn_metadata.tile_elems == 64 and sm100a_requested)
    needs_sm100a_pair = n_tiles % 2 != 0 and not grad_mode and sm100a_route
    kernel_tiles = n_tiles + int(needs_sm100a_pair)
    target_shape = (x.shape[0], kernel_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])

    # A grad-tracking forward must not reuse the builder-owned buffer. The
    # holder outlives the training step -- one builder serves the whole run,
    # and every step's metadata references the same holder -- while every
    # VSA layer writes this one buffer in place. A graph-tracked tiled
    # tensor left on the holder therefore anchors the step's in-place
    # autograd edges, and through them the activations they saved, for the
    # rest of training. This is the VSA-H3 counterpart of the Wan tile-cache
    # OOM (#1423), which training fixed with ``vsa_cache_tile_buf=False``.
    if grad_mode:
        return scatter_into_tile_buf(x, target_shape, attn_metadata.untile_combined_index, None)

    # ``untile_combined_index`` maps each packed row to a logical tile
    # slot. Different geometries can share one transport shape; clear a
    # reused allocation once when the mapping identity changes so no old
    # valid row can survive as padding.
    holder = attn_metadata.tile_buf_holder
    if holder is None:
        raise RuntimeError("VSA-H3 metadata has no builder-owned tile buffer holder")
    if (not compiling and attn_metadata.tile_elems in _SM100A_TILE_ELEMS
            and envs.FASTVIDEO_H3_VSA_HEADS_FIRST_TILE.get()
            and supports_heads_first_scatter(x, attn_metadata.untile_combined_index)):
        return self._tile_heads_first(x, attn_metadata, holder, kernel_tiles, needs_sm100a_pair)
    buffer_matches = (holder.buffer is not None and holder.buffer.shape == target_shape
                      and holder.buffer.dtype == x.dtype and holder.buffer.device == x.device)
    if buffer_matches and holder.untile_geometry is not attn_metadata.untile_combined_index:
        holder.buffer.zero_()
    holder.buffer = scatter_into_tile_buf(x, target_shape, attn_metadata.untile_combined_index, holder.buffer)
    holder.untile_geometry = attn_metadata.untile_combined_index
    if needs_sm100a_pair:
        # A prior even geometry can reuse this allocation and may have
        # written the last tile as logical data.
        holder.buffer[:, n_tiles * attn_metadata.tile_elems:].zero_()
    return holder.buffer

fastvideo.attention.backends.video_sparse_attn_h3.MiniMaxH3VSAMetadataBuilder

MiniMaxH3VSAMetadataBuilder()

Bases: AttentionMetadataBuilder

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def __init__(self) -> None:
    self._tile_buf_holder = _MiniMaxH3VSATileBufferHolder()

Methods:

fastvideo.attention.backends.video_sparse_attn_h3.MiniMaxH3VSAMetadataBuilder.build
build(current_timestep: int, patch_size: tuple[int, int, int], VSA_sparsity: float, packed_segments: tuple[int | tuple[int, int, int], ...], device: device, exempt: bool = True, dense_layers: tuple[int, ...] = (), tile_size: int = _TILE_ELEMS, ref_keep_rate: float | None = None, **kwargs: dict[str, Any]) -> MiniMaxH3VSAMetadata

Build per-step metadata for one packed H3 sequence.

packed_segments lists the sequence in packed order: an int is a dense segment's row count, and a (t, h, w) triple is the raw latent shape of one sparse video region. The last region is the generated video and follows VSA_sparsity; every earlier region is a reference video and keeps ref_keep_rate of its tiles, which a build with reference-video regions requires. exempt=False (compete mode) supports only the generated-video region.

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def build(  # type: ignore
    self,
    current_timestep: int,
    patch_size: tuple[int, int, int],
    VSA_sparsity: float,
    packed_segments: tuple[int | tuple[int, int, int], ...],
    device: torch.device,
    exempt: bool = True,
    dense_layers: tuple[int, ...] = (),
    tile_size: int = _TILE_ELEMS,
    ref_keep_rate: float | None = None,
    **kwargs: dict[str, Any],
) -> MiniMaxH3VSAMetadata:
    """Build per-step metadata for one packed H3 sequence.

    ``packed_segments`` lists the sequence in packed order: an ``int`` is
    a dense segment's row count, and a ``(t, h, w)`` triple is the raw
    latent shape of one sparse video region. The last region is the
    generated video and follows ``VSA_sparsity``; every earlier region is
    a reference video and keeps ``ref_keep_rate`` of its tiles, which a
    build with reference-video regions requires. ``exempt=False``
    (compete mode) supports only the generated-video region.
    """
    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}")
    # Video regions become token grids under the patch size; empty dense segments are dropped.
    token_segments: list[int | tuple[int, int, int]] = []
    for segment in packed_segments:
        if isinstance(segment, tuple):
            if any(int(v) <= 0 or int(v) % p for v, p in zip(segment, patch_size, strict=True)):
                raise ValueError(f"VSA-H3 video region latent shape {tuple(segment)} is not a positive multiple "
                                 f"of patch {tuple(patch_size)}.")
            t, h, w = (int(v) // p for v, p in zip(segment, patch_size, strict=True))
            token_segments.append((t, h, w))
        elif segment > 0:
            token_segments.append(int(segment))
    num_regions = sum(isinstance(segment, tuple) for segment in token_segments)
    if num_regions == 0:
        raise ValueError("VSA-H3 needs at least one video region.")
    reference_sparsities: tuple[float, ...] = ()
    if num_regions > 1:
        if ref_keep_rate is None:
            raise ValueError(f"A VSA-H3 build with {num_regions - 1} reference-video region(s) needs "
                             "ref_keep_rate.")
        if not exempt:
            raise ValueError(f"vsa_mode='compete' supports only the generated-video region; this build has "
                             f"{num_regions - 1} reference-video region(s). Use vsa_mode='exempt'.")
        reference_sparsities = (1.0 - float(ref_keep_rate), ) * (num_regions - 1)
    total_seq_length = sum(math.prod(s) if isinstance(s, tuple) else s for s in token_segments)

    (_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles,
     video_tile_spans) = _h3_segment_tile_geometry(tuple(token_segments), device, tile_shape)

    # A dense build (sparsity 0: dense steps) keeps every region dense,
    # whatever the reference keep rate.
    VSA_sparsity = float(VSA_sparsity)
    span_sparsities = (0.0, ) * num_regions if VSA_sparsity <= 0.0 else (*reference_sparsities, VSA_sparsity)

    dense_layers = tuple(int(layer) for layer in dense_layers)
    return MiniMaxH3VSAMetadata(
        current_timestep=current_timestep,
        VSA_sparsity=VSA_sparsity,
        total_seq_length=total_seq_length,
        num_prefix_tiles=num_prefix_tiles,
        num_video_tiles=num_video_tiles,
        exempt=exempt,
        variable_block_sizes=variable_block_sizes,
        untile_combined_index=untile_combined_index,
        dense_layers_tensor=torch.tensor(dense_layers, device=device, dtype=torch.int64),
        video_tile_spans=video_tile_spans,
        span_sparsities=span_sparsities,
        tile_elems=int(tile_size),
        dense_layers=dense_layers,
        tile_buf_holder=self._tile_buf_holder,
    )

Functions:

fastvideo.attention.backends.video_sparse_attn_h3.token_tile_and_valid

token_tile_and_valid(variable_block_sizes: Tensor, tile_elems: int = _TILE_ELEMS) -> tuple[Tensor, Tensor]

Per padded-token tile id and pad-validity mask.

The single encoding of the padding contract, shared by the probe and the test oracle so they cannot drift from the backend's tile geometry. tile_elems must match the metadata the sizes came from (MiniMaxH3VSAMetadata.tile_elems).

Source code in fastvideo/attention/backends/video_sparse_attn_h3.py
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
                         tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
    """Per padded-token tile id and pad-validity mask.

    The single encoding of the padding contract, shared by the probe and the
    test oracle so they cannot drift from the backend's tile geometry.
    ``tile_elems`` must match the metadata the sizes came from
    (``MiniMaxH3VSAMetadata.tile_elems``).
    """
    device = variable_block_sizes.device
    token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
    token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
    return token_tile, token_valid