| """VSA-H3 — MiniMax-H3's Video Sparse Attention, as a `diffusers` attention processor. |
| |
| `FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree` is distilled **under** block-sparse attention |
| (`attention_backend: VIDEO_SPARSE_ATTN_H3`, `vsa_tile_size: 64`, `vsa_sparsity: 0.9` in the checkpoint's own |
| `fastvideo_inference.json`), and FastVideo's release notes are explicit that dense attention is *not* a drop-in |
| substitute for a VSA-trained student: the student learned to attend to the top-10% tiles its selector picks, and it |
| also ships 50 trained `attn.to_gate_compress` tensors that only the sparse path reads. So this Space runs the sparse |
| path. |
| |
| `diffusers` has no VSA backend, so this module ports FastVideo's `MiniMaxH3VSABackend` |
| (`fastvideo/attention/backends/video_sparse_attn_h3.py`) onto `MiniMaxH3Attention`: |
| |
| 1. **Tiling.** The packed sequence `[text | audio | video]` is cut into 64-token tiles: segment-pure prefix chunks |
| first, then `(4, 4, 4)` cubes of the post-patchify `(t, h, w)` video grid. Tiles are zero-padded to 64 and |
| `variable_block_sizes` carries each tile's true occupancy. |
| 2. **Selection.** Per head, tiles are mean-pooled in fp32, `scores = q_pooled @ k_pooledᵀ / √d`, and each query tile |
| keeps the top `ceil((1 - 0.9) * num_video_tiles)` video tiles. Prefix (text/audio) keys are *exempt* — always |
| selected — and prefix queries are always dense, which is FastVideo's default `vsa_mode`. |
| 3. **Kernel.** The resulting bool block map is compacted with FastVideo's `map_to_index` Triton kernel and consumed by |
| its `triton_block_sparse_attn_forward`, both vendored verbatim under `vsa_kernel/`. This is FastVideo's own |
| `--vsa-kernel triton` route; the checkpoint's `vsa_kernel: sm100a` is the GB200-only fast path for the *same* mask |
| semantics, and this Space's Blackwell RTX PRO 6000 is sm120. |
| 4. **Compression branch.** `out_c = softmax(scores) @ v_pooled` broadcast back over each tile's rows and scaled by the |
| trained per-row `to_gate_compress` gate, added to the sparse output. The base MiniMax-H3 release zero-initializes |
| this gate (branch inert); this student ships it trained. |
| |
| The one deliberate deviation from the reference is memory layout, not math: buffers are tiled straight into the |
| kernel's `[B, H, S_pad, D]` layout and reused across the 50 blocks, and the gate is applied *after* untiling, so the |
| sparse path costs one persistent tile buffer per projection instead of four plus three transposed copies. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import functools |
| import math |
| import os |
| from dataclasses import dataclass |
|
|
| import torch |
|
|
| |
| |
| TILE_ELEMS = 64 |
| TILE_SHAPE = (4, 4, 4) |
| DEFAULT_SPARSITY = 0.9 |
|
|
|
|
| |
| |
| |
|
|
|
|
| def _tile_partition_indices(dit_seq_shape: tuple[int, int, int], device: torch.device) -> torch.Tensor: |
| """Row indices of the video grid in `(4, 4, 4)` tile order.""" |
| grid_t, grid_h, grid_w = dit_seq_shape |
| ts, hs, ws = TILE_SHAPE |
| indices = torch.arange(grid_t * grid_h * grid_w, device=device, dtype=torch.long).reshape(grid_t, grid_h, grid_w) |
| chunks = [] |
| for t in range(math.ceil(grid_t / ts)): |
| for h in range(math.ceil(grid_h / hs)): |
| for w in range(math.ceil(grid_w / ws)): |
| chunks.append( |
| indices[t * ts : min(t * ts + ts, grid_t), h * hs : min(h * hs + hs, grid_h), w * ws : min(w * ws + ws, grid_w)].flatten() |
| ) |
| return torch.cat(chunks, dim=0) |
|
|
|
|
| def _video_block_sizes(dit_seq_shape: tuple[int, int, int], device: torch.device) -> torch.Tensor: |
| """Valid (non-padded) token count of every video tile, in the same tile order.""" |
| ts, hs, ws = TILE_SHAPE |
|
|
| def sizes(length: int, tile: int) -> torch.Tensor: |
| count = math.ceil(length / tile) |
| out = torch.full((count,), tile, dtype=torch.long, device=device) |
| remainder = length - (count - 1) * tile |
| out[-1] = remainder if remainder > 0 else tile |
| return out |
|
|
| t_sizes, h_sizes, w_sizes = (sizes(length, tile) for length, tile in zip(dit_seq_shape, TILE_SHAPE)) |
| return (t_sizes[:, None, None] * h_sizes[None, :, None] * w_sizes[None, None, :]).reshape(-1) |
|
|
|
|
| def _non_pad_index(variable_block_sizes: torch.Tensor) -> torch.Tensor: |
| """Padded-buffer slot of every real token, in tile order.""" |
| device = variable_block_sizes.device |
| starts = torch.arange(variable_block_sizes.shape[0], device=device) * TILE_ELEMS |
| slots = starts[:, None] + torch.arange(TILE_ELEMS, device=device)[None, :] |
| keep = torch.arange(TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None] |
| return slots[keep] |
|
|
|
|
| @dataclass(frozen=True, eq=False) |
| class VSAGeometry: |
| """Everything the sparse path needs about one packed layout. Cached per `(prefix_segments, grid)`.""" |
|
|
| seq_len: int |
| padded_len: int |
| n_tiles: int |
| num_prefix_tiles: int |
| num_video_tiles: int |
| topk: int |
| |
| variable_block_sizes: torch.Tensor |
| |
| tile_divisor: torch.Tensor |
| |
| untile_index: torch.Tensor |
| |
| row_tile_index: torch.Tensor |
|
|
| @property |
| def dense(self) -> bool: |
| return self.topk >= self.num_video_tiles |
|
|
|
|
| @functools.lru_cache(maxsize=8) |
| def build_geometry( |
| prefix_segments: tuple[int, ...], |
| dit_seq_shape: tuple[int, int, int], |
| device: torch.device, |
| sparsity: float = DEFAULT_SPARSITY, |
| ) -> VSAGeometry: |
| prefix_len = sum(prefix_segments) |
|
|
| prefix_sizes: list[int] = [] |
| for segment in prefix_segments: |
| full, remainder = divmod(segment, TILE_ELEMS) |
| prefix_sizes.extend([TILE_ELEMS] * full) |
| if remainder: |
| prefix_sizes.append(remainder) |
|
|
| video_sizes = _video_block_sizes(dit_seq_shape, device) |
| variable_block_sizes = torch.cat( |
| [torch.tensor(prefix_sizes, dtype=torch.long, device=device), video_sizes] |
| ) |
| partition = torch.cat( |
| [ |
| torch.arange(prefix_len, device=device, dtype=torch.long), |
| _tile_partition_indices(dit_seq_shape, device) + prefix_len, |
| ] |
| ) |
| untile_index = _non_pad_index(variable_block_sizes)[torch.argsort(partition)] |
|
|
| n_tiles = int(variable_block_sizes.numel()) |
| num_video_tiles = int(video_sizes.numel()) |
| return VSAGeometry( |
| seq_len=int(partition.numel()), |
| padded_len=n_tiles * TILE_ELEMS, |
| n_tiles=n_tiles, |
| num_prefix_tiles=len(prefix_sizes), |
| num_video_tiles=num_video_tiles, |
| |
| topk=max(1, min(math.ceil((1.0 - sparsity) * num_video_tiles), num_video_tiles)), |
| variable_block_sizes=variable_block_sizes.to(torch.int32).contiguous(), |
| tile_divisor=variable_block_sizes.to(torch.float32).view(1, 1, -1, 1), |
| untile_index=untile_index, |
| row_tile_index=untile_index // TILE_ELEMS, |
| ) |
|
|
|
|
| |
| |
| |
| |
| |
| _REQUEST_STATE: dict = {"nonce": 0, "geometry": None} |
| _REQUEST_COUNTER = 0 |
|
|
|
|
| def begin_request() -> None: |
| """Open a new packed-sequence request: drop the reused geometry so the next forward re-derives it.""" |
| global _REQUEST_COUNTER |
|
|
| _REQUEST_COUNTER += 1 |
| _REQUEST_STATE["nonce"] = _REQUEST_COUNTER |
| _REQUEST_STATE["geometry"] = None |
|
|
|
|
| def geometry_from_layout( |
| token_tags: torch.Tensor, |
| position_ids: torch.Tensor, |
| sparsity: float = DEFAULT_SPARSITY, |
| ) -> VSAGeometry | None: |
| """Recover the VSA geometry from what the transformer is actually given. |
| |
| The packed sequence a `t2va` request builds is `[text | audio | video]`, but nothing downstream is told that, so |
| the layout is read back off the two per-row descriptions the transformer already takes: `token_tags` (0 video, |
| 1 text, 2 audio) gives the segment boundaries, and the `(t, h, w)` rotary grid of the video rows gives the shape |
| of the video block. Returns `None` for any layout the sparse path does not cover, so the caller can stay dense. |
| """ |
| state = _REQUEST_STATE |
| cached = state["geometry"] |
| if cached is not None and state["nonce"] == _REQUEST_COUNTER and cached.seq_len == token_tags.shape[0]: |
| return cached |
|
|
| tags = token_tags.tolist() |
| seq_len = len(tags) |
| if seq_len == 0 or tags[-1] != 0: |
| return None |
|
|
| video_start = seq_len |
| while video_start > 0 and tags[video_start - 1] == 0: |
| video_start -= 1 |
| if video_start == 0: |
| return None |
|
|
| prefix_segments: list[int] = [] |
| previous = None |
| for tag in tags[:video_start]: |
| if tag == previous: |
| prefix_segments[-1] += 1 |
| else: |
| prefix_segments.append(1) |
| previous = tag |
|
|
| |
| |
| grid = position_ids[video_start:] |
| rows_per_frame = int((grid[:, 0] == grid[0, 0]).sum()) |
| grid_w = int((grid[:rows_per_frame, 1] == grid[0, 1]).sum()) |
| if rows_per_frame == 0 or grid_w == 0 or rows_per_frame % grid_w: |
| return None |
| grid_h = rows_per_frame // grid_w |
| num_video_rows = seq_len - video_start |
| if num_video_rows % rows_per_frame: |
| return None |
| grid_t = num_video_rows // rows_per_frame |
|
|
| geometry = build_geometry(tuple(prefix_segments), (grid_t, grid_h, grid_w), token_tags.device, sparsity) |
| state["geometry"] = geometry |
| return geometry |
|
|
|
|
| |
| |
| |
|
|
| |
| |
| _TILE_BUFFERS: dict[str, tuple[torch.Tensor, int]] = {} |
|
|
|
|
| def reset_tile_buffers() -> None: |
| _TILE_BUFFERS.clear() |
|
|
|
|
| def _tile(x: torch.Tensor, geometry: VSAGeometry, slot: str) -> torch.Tensor: |
| """`[B, S, H, D]` -> the kernel's `[B, H, S_pad, D]`, pad slots zero.""" |
| batch, _, heads, dim = x.shape |
| shape = (batch, heads, geometry.padded_len, dim) |
| cached = _TILE_BUFFERS.get(slot) |
| if cached is None or cached[0].shape != shape or cached[0].dtype != x.dtype or cached[0].device != x.device: |
| buffer = torch.zeros(shape, dtype=x.dtype, device=x.device) |
| else: |
| buffer = cached[0] |
| if cached[1] != id(geometry): |
| buffer.zero_() |
| buffer.index_copy_(2, geometry.untile_index, x.transpose(1, 2)) |
| _TILE_BUFFERS[slot] = (buffer, id(geometry)) |
| return buffer |
|
|
|
|
| def _pool(tiled: torch.Tensor, geometry: VSAGeometry) -> torch.Tensor: |
| """fp32 masked mean over each 64-token tile. `[B, H, S_pad, D]` -> `[B, H, n_tiles, D]`.""" |
| batch, heads, _, dim = tiled.shape |
| pooled = tiled.view(batch, heads, geometry.n_tiles, TILE_ELEMS, dim).sum(dim=3, dtype=torch.float32) |
| return pooled / geometry.tile_divisor |
|
|
|
|
| def _block_mask(scores: torch.Tensor, geometry: VSAGeometry) -> torch.Tensor: |
| """Top-k video tiles per query tile, with the prefix exempt and prefix queries dense.""" |
| prefix = geometry.num_prefix_tiles |
| mask = torch.zeros_like(scores, dtype=torch.bool) |
| |
| |
| indices = scores[..., prefix:].topk(geometry.topk, dim=-1, sorted=False).indices + prefix |
| mask.scatter_(-1, indices, True) |
| mask[..., :prefix] = True |
| mask[:, :, :prefix, :] = True |
| return mask |
|
|
|
|
| |
| |
| |
| |
| |
| |
| _CUDA_SPARSE: dict = {"resolved": False, "op": None} |
|
|
|
|
| def _resolve_cuda_sparse_op(): |
| if _CUDA_SPARSE["resolved"]: |
| return _CUDA_SPARSE["op"] |
| _CUDA_SPARSE["resolved"] = True |
| if os.environ.get("H3_VSA_CUDA", "1") != "1": |
| return None |
| try: |
| from fastvideo_kernel import block_sparse_attn_sm100a as sm100a |
|
|
| if not sm100a._HAS_VSA_SM100A: |
| |
| from fastvideo_kernel._C import fastvideo_kernel_ops as _ops |
|
|
| available = sorted(a for a in dir(_ops) if "sparse" in a.lower()) |
| raise ImportError(f"sm100a forward flag unset; pybind exposes {available}") |
| op = sm100a.block_sparse_attn_sm100a |
| _CUDA_SPARSE["op"] = op |
| print("[vsa] CUDA sparse-attention kernel: fastvideo_kernel (sm100a fatbin on sm120)", flush=True) |
| except Exception as error: |
| print(f"[vsa] fastvideo_kernel unavailable ({type(error).__name__}: {error}); using vendored Triton", flush=True) |
| _CUDA_SPARSE["op"] = None |
| return _CUDA_SPARSE["op"] |
|
|
|
|
| def sparse_attention( |
| query: torch.Tensor, |
| key: torch.Tensor, |
| value: torch.Tensor, |
| gate_compress: torch.Tensor | None, |
| geometry: VSAGeometry, |
| ) -> torch.Tensor: |
| """VSA-H3 over one packed sequence. All tensors are `[B, S, H, D]`; the result is too.""" |
| from vsa_kernel import map_to_index, triton_block_sparse_attn_forward |
|
|
| query_tiled = _tile(query, geometry, "q") |
| key_tiled = _tile(key, geometry, "k") |
| value_tiled = _tile(value, geometry, "v") |
|
|
| scores = torch.matmul(_pool(query_tiled, geometry), _pool(key_tiled, geometry).transpose(-2, -1)) |
| scores = scores / math.sqrt(query.shape[-1]) |
|
|
| mask = _block_mask(scores, geometry) |
| q2k_index, q2k_num = map_to_index(mask) |
| |
| |
| out_tiled = None |
| cuda_op = _resolve_cuda_sparse_op() |
| if cuda_op is not None and geometry.n_tiles % 2 == 0: |
| try: |
| out_tiled, _ = cuda_op( |
| query_tiled, |
| key_tiled, |
| value_tiled, |
| q2k_index, |
| q2k_num, |
| geometry.variable_block_sizes, |
| need_lse=False, |
| ) |
| except Exception as error: |
| print(f"[vsa] CUDA kernel failed ({type(error).__name__}: {error}); falling back to Triton", flush=True) |
| _CUDA_SPARSE["op"] = None |
| out_tiled = None |
| if out_tiled is None: |
| out_tiled, _ = triton_block_sparse_attn_forward( |
| query_tiled, |
| key_tiled, |
| value_tiled, |
| q2k_index, |
| q2k_num, |
| geometry.variable_block_sizes, |
| ) |
| out = out_tiled.index_select(2, geometry.untile_index) |
|
|
| if gate_compress is not None: |
| |
| |
| pooled = torch.matmul(torch.softmax(scores, dim=-1), _pool(value_tiled, geometry)).to(out.dtype) |
| contribution = pooled.index_select(2, geometry.row_tile_index) |
| del pooled |
| contribution.mul_(gate_compress.transpose(1, 2)) |
| out.add_(contribution) |
| del contribution |
|
|
| return out.transpose(1, 2) |
|
|
|
|
| |
| |
| |
|
|
| |
| |
| _ACTIVE: dict[str, VSAGeometry | None] = {"geometry": None} |
|
|
| |
| |
| |
| LAST_LAYOUT = "no forward yet" |
| _REPORTED: set[str] = set() |
|
|
|
|
| def _report(summary: str) -> None: |
| global LAST_LAYOUT |
|
|
| LAST_LAYOUT = summary |
| if summary not in _REPORTED: |
| _REPORTED.add(summary) |
| print(f"[vsa] {summary}", flush=True) |
|
|
|
|
| class MiniMaxH3VSAAttnProcessor: |
| """`MiniMaxH3AttnProcessor` with `dispatch_attention_fn` replaced by VSA-H3.""" |
|
|
| _attention_backend = None |
| _parallel_config = None |
|
|
| def __call__(self, attn, hidden_states, rotary_emb=None, attention_mask=None): |
| from diffusers.models.transformers.transformer_minimax_h3 import _apply_rotary_emb |
|
|
| geometry = _ACTIVE["geometry"] |
|
|
| query = attn.to_q(hidden_states).unflatten(-1, (attn.heads, -1)) |
| key = attn.to_k(hidden_states).unflatten(-1, (attn.heads, -1)) |
| value = attn.to_v(hidden_states).unflatten(-1, (attn.heads, -1)) |
| query = attn.norm_q(query) |
| key = attn.norm_k(key) |
| if rotary_emb is not None: |
| query = _apply_rotary_emb(query, *rotary_emb) |
| key = _apply_rotary_emb(key, *rotary_emb) |
|
|
| if geometry is None or geometry.seq_len != hidden_states.shape[1]: |
| _report(f"DENSE FALLBACK in the attention processor at seq_len={hidden_states.shape[1]}") |
| from diffusers.models.attention_dispatch import dispatch_attention_fn |
|
|
| out = dispatch_attention_fn( |
| query, |
| key, |
| value, |
| attn_mask=attention_mask, |
| dropout_p=0.0, |
| is_causal=False, |
| backend=self._attention_backend, |
| parallel_config=self._parallel_config, |
| ) |
| else: |
| gate_compress = None |
| gate = getattr(attn, "to_gate_compress", None) |
| if gate is not None: |
| gate_compress = gate(hidden_states).unflatten(-1, (attn.heads, -1)) |
| out = sparse_attention(query, key, value, gate_compress, geometry) |
|
|
| out = out.flatten(2, 3).type_as(query) |
| out = attn.to_out[0](out) |
| return attn.to_out[1](out) |
|
|
|
|
| def add_gate_compress_modules() -> None: |
| """Give every block's attention the `to_gate_compress` projection the base `diffusers` port has no use for. |
| |
| The trained gate lives in the checkpoint as 50 `transformer_blocks.*.attn.to_gate_compress.weight` tensors, but |
| `MiniMaxH3Attention` does not declare the module, so `from_pretrained` reports them as unexpected and drops them. |
| Declaring it before the transformer is instantiated is what makes them load. |
| """ |
| from diffusers.models.transformers import transformer_minimax_h3 as module |
|
|
| block_cls = module.MiniMaxH3TransformerBlock |
| if getattr(block_cls, "_vsa_gate_patched", False): |
| return |
| original_init = block_cls.__init__ |
|
|
| def patched_init(self, hidden_size, num_attention_heads, attention_head_dim, *args, **kwargs): |
| original_init(self, hidden_size, num_attention_heads, attention_head_dim, *args, **kwargs) |
| self.attn.to_gate_compress = torch.nn.Linear( |
| hidden_size, num_attention_heads * attention_head_dim, bias=False |
| ) |
|
|
| block_cls.__init__ = patched_init |
| block_cls._vsa_gate_patched = True |
|
|
|
|
| def install(transformer, sparsity: float = DEFAULT_SPARSITY) -> tuple[int, int]: |
| """Put the VSA processor on the 50 packed-sequence blocks and wrap `forward` to publish the layout. |
| |
| The token refiner has its own block class, over the *text* stream rather than the packed sequence, so it is |
| untouched. Returns `(blocks, live gates)`. |
| """ |
| installed = 0 |
| gates = 0 |
| for block in transformer.transformer_blocks: |
| gate = getattr(block.attn, "to_gate_compress", None) |
| if gate is not None and not bool((gate.weight != 0).any()): |
| |
| |
| block.attn.to_gate_compress = None |
| elif gate is not None: |
| gates += 1 |
| block.attn.set_processor(MiniMaxH3VSAAttnProcessor()) |
| installed += 1 |
|
|
| if getattr(transformer, "_vsa_forward_wrapped", False): |
| return installed, gates |
| original_forward = transformer.forward |
|
|
| @functools.wraps(original_forward) |
| def forward(*args, **kwargs): |
| token_tags = kwargs.get("token_tags") |
| position_ids = kwargs.get("position_ids") |
| if token_tags is None or position_ids is None: |
| |
| import inspect |
|
|
| bound = inspect.signature(original_forward).bind_partial(*args, **kwargs).arguments |
| token_tags = token_tags if token_tags is not None else bound.get("token_tags") |
| position_ids = position_ids if position_ids is not None else bound.get("position_ids") |
| geometry = None |
| if token_tags is not None and position_ids is not None: |
| geometry = geometry_from_layout(token_tags, position_ids, sparsity) |
| if geometry is None: |
| _report("DENSE FALLBACK: no VSA geometry could be derived from this forward's layout") |
| else: |
| _report( |
| f"sparse: seq={geometry.seq_len} padded={geometry.padded_len} tiles={geometry.n_tiles} " |
| f"(prefix {geometry.num_prefix_tiles} + video {geometry.num_video_tiles}) " |
| f"topk={geometry.topk} ({geometry.topk / geometry.num_video_tiles:.1%} of video tiles kept)" |
| ) |
| _ACTIVE["geometry"] = geometry |
| try: |
| return original_forward(*args, **kwargs) |
| finally: |
| _ACTIVE["geometry"] = None |
|
|
| transformer.forward = forward |
| transformer._vsa_forward_wrapped = True |
| return installed, gates |
|
|