Add initial support for Minimax H3 VSA sparse attention

This commit is contained in:
Jaret Burkett
2026-08-30 08:24:43 -06:00
parent be995185f5
commit 2a69c1e7de
11 changed files with 1894 additions and 18 deletions

View File

@@ -19,7 +19,7 @@ from .prx_pixel_t2i import PRXPixelT2IModel
from .krea2 import Krea2Model from .krea2 import Krea2Model
from .boogu_image import BooguImageModel, BooguImageEditModel from .boogu_image import BooguImageModel, BooguImageEditModel
from .mageflow import MageFlowModel, MageFlowEditModel from .mageflow import MageFlowModel, MageFlowEditModel
from .minimax_h3 import MinimaxH3Model, MinimaxH3Ref2VAModel from .minimax_h3 import MinimaxH3Model, MinimaxH3Ref2VAModel, MinimaxH3FastModel
AI_TOOLKIT_MODELS = [ AI_TOOLKIT_MODELS = [
# put a list of models here # put a list of models here
@@ -58,4 +58,5 @@ AI_TOOLKIT_MODELS = [
MageFlowEditModel, MageFlowEditModel,
MinimaxH3Model, MinimaxH3Model,
MinimaxH3Ref2VAModel, MinimaxH3Ref2VAModel,
MinimaxH3FastModel,
] ]

View File

@@ -48,6 +48,7 @@ try:
convert_ltx2_audio_vae, convert_ltx2_audio_vae,
convert_ltx2_vocoder, convert_ltx2_vocoder,
convert_ltx2_connectors, convert_ltx2_connectors,
split_transformer_and_connector_state_dict,
dequantize_state_dict, dequantize_state_dict,
convert_comfy_gemma3_to_transformers, convert_comfy_gemma3_to_transformers,
convert_lora_original_to_diffusers, convert_lora_original_to_diffusers,
@@ -294,6 +295,16 @@ class LTX2Model(BaseModel):
original_dit_ckpt, version=self.ltx_version original_dit_ckpt, version=self.ltx_version
) )
transformer = transformer.to(dtype) transformer = transformer.to(dtype)
# the transformer holds these tensors (assign=True); drop the dict refs
# so each block's bf16 original frees as it quantizes instead of the
# whole dit staying in RAM. Connector keys stay — converted later.
transformer_sd, _ = split_transformer_and_connector_state_dict(
original_dit_ckpt
)
for key in transformer_sd:
combined_state_dict.pop(dit_prefix + key, None)
del transformer_sd, original_dit_ckpt
flush()
else: else:
if os.path.exists(model_path): if os.path.exists(model_path):
# check if the path is a full checkpoint. # check if the path is a full checkpoint.
@@ -1322,6 +1333,13 @@ class LTX25Model(LTX2Model):
transformer, transformer_sd, "transformer" transformer, transformer_sd, "transformer"
) )
del transformer_sd del transformer_sd
# dit_sd still references every transformer tensor (assign=True sharing);
# drop them so each block frees as it quantizes. Connector keys stay for
# the convert_ltx2_connectors call below.
trans_sd, _ = split_transformer_and_connector_state_dict(dit_sd)
for key in trans_sd:
dit_sd.pop(key, None)
del trans_sd
if num_quantized_dit == 0: if num_quantized_dit == 0:
transformer = transformer.to(dtype) transformer = transformer.to(dtype)
else: else:

View File

@@ -1 +1 @@
from .minimax_h3 import MinimaxH3Model, MinimaxH3Ref2VAModel from .minimax_h3 import MinimaxH3Model, MinimaxH3Ref2VAModel, MinimaxH3FastModel

View File

@@ -118,7 +118,15 @@ COMFY_FILES = {
"text_encoder": "text_encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", "text_encoder": "text_encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors",
"video_vae": "vae/minimax_h3_video_vae_fp16.safetensors", "video_vae": "vae/minimax_h3_video_vae_fp16.safetensors",
"audio_vae": "vae/minimax_h3_audio_vae_fp32.safetensors", "audio_vae": "vae/minimax_h3_audio_vae_fp32.safetensors",
"dit_fasth3": "diffusion_models/minimax_h3_fasth3_preview_v0.2_int8_convrot.safetensors",
} }
# FastH3 (FastVideo 4-step VSA distill) int8-convrot repack, produced by
# scripts/convert_minimax_h2_fastvideo.py from the FastVideo diffusers repo.
# Hub fallback repo for the file (flat at the repo root, downloaded into
# MODELS_PATH/diffusion_models/); not published there yet — until it is, the
# file must exist locally or be built with the converter.
FASTH3_REPO = "Kijai/MiniMax-H3-experimental"
FASTH3_SOURCE_REPO = "FastVideo/FastVideo-Minimax-FastH3-Preview-v0.2"
# tokenizer/processor/text-encoder config come from the original repo (tiny files) # tokenizer/processor/text-encoder config come from the original repo (tiny files)
ORIGINAL_REPO = "MiniMaxAI/MiniMax-H3" ORIGINAL_REPO = "MiniMaxAI/MiniMax-H3"
@@ -248,9 +256,7 @@ class MinimaxH3Model(BaseModel):
return resolve_comfy_file( return resolve_comfy_file(
COMFY_FILES[component], COMFY_FILES[component],
repo_id=repo_id_from_name_or_path(name_or_path, COMFY_REPO), repo_id=repo_id_from_name_or_path(name_or_path, COMFY_REPO),
override_path=self.model_config.model_kwargs.get( override_path=self.model_config.model_kwargs.get(f"{component}_path", None),
f"{component}_path", None
),
extra_roots=extra_roots, extra_roots=extra_roots,
status_fn=self.print_and_status_update, status_fn=self.print_and_status_update,
) )
@@ -284,9 +290,7 @@ class MinimaxH3Model(BaseModel):
lora_path = self.model_config.assistant_lora_path lora_path = self.model_config.assistant_lora_path
if not os.path.exists(lora_path): if not os.path.exists(lora_path):
filename = os.path.basename(lora_path) filename = os.path.basename(lora_path)
found = find_file_recursive( found = find_file_recursive(os.path.join(MODELS_PATH, "loras"), filename)
os.path.join(MODELS_PATH, "loras"), filename
)
if found is not None: if found is not None:
lora_path = found lora_path = found
else: else:
@@ -955,6 +959,8 @@ class MinimaxH3Model(BaseModel):
video_indices=video_indices.to(device), video_indices=video_indices.to(device),
audio_indices=audio_indices.to(device), audio_indices=audio_indices.to(device),
text_indices=text_indices.to(device), text_indices=text_indices.to(device),
# target-video token grid (patch 1x2x2); consumed only by VSA models
vsa_video_grid=(t_lat, h_lat // 2, w_lat // 2),
) )
if num_cond_audio > 0: if num_cond_audio > 0:
@@ -1091,6 +1097,7 @@ class MinimaxH3Model(BaseModel):
# standard diffusion_model prefix maps directly # standard diffusion_model prefix maps directly
lora_keys_use_comfy_prefix = True lora_keys_use_comfy_prefix = True
class MinimaxH3Ref2VAModel(MinimaxH3Model): class MinimaxH3Ref2VAModel(MinimaxH3Model):
"""Reference-to-video (ref2va): the control images ride along as reference """Reference-to-video (ref2va): the control images ride along as reference
blocks in the packed sequence (plus ``<Picture i>: `` vision blocks in the blocks in the packed sequence (plus ``<Picture i>: `` vision blocks in the
@@ -1548,3 +1555,109 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
if is_video: if is_video:
return result # dict consumed by new_save_image_function return result # dict consumed by new_save_image_function
return result[0] return result[0]
class MinimaxH3FastModel(MinimaxH3Model):
"""FastH3: FastVideo's DMD2-distilled 4-step MiniMax-H3 preview, trained
with VSA (Video Sparse Attention) at 90% sparsity. Text-to-video(+audio)
only — no first-frame or reference conditioning.
The DiT is the int8-ConvRot repack of
``FastVideo/FastVideo-Minimax-FastH3-Preview-v0.2`` (built with
scripts/convert_minimax_h2_fastvideo.py): the pruned comfy layout plus a
``blocks.N.attn.to_gate_compress`` linear per block. Attention runs the
trained VSA policy — pooled 64-token (4,4,4) tile top-k block-sparse
attention with the gated compression branch — via FlexAttention
(src/vsa.py), in training and sampling alike. Text encoder and VAEs are
shared with the base model.
``model_kwargs``:
- ``vsa`` (default true): false runs dense attention with the gate
branch off, matching FastVideo's dense LoRA-preview mode
- ``vsa_sparsity`` (default 0.9, the checkpoint's trained policy)
Sample with 4 steps (the distilled schedule) and guidance_scale 1.
"""
arch = "minimax_h3_vsa"
# FastH3 samples on its trained ladder: [999, 749, 500, 250] on the
# shared 1000-step grid, each scheduler applying its own shift
t1000_sample_ladder = True
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
kw = self.model_config.model_kwargs
self.vsa_sparsity: Optional[float] = None
if bool(kw.get("vsa", True)):
self.vsa_sparsity = float(kw.get("vsa_sparsity", 0.9))
def _dit_component(self) -> str:
return "dit_fasth3"
def _resolve_comfy_file(self, component: str) -> str:
if component != "dit_fasth3":
return super()._resolve_comfy_file(component)
# local search at the comfy-layout locations first; the hub file sits
# at the FastH3 repo's root, so the download is done directly here
name_or_path = self.model_config.name_or_path
extra_roots = (
[name_or_path] if name_or_path and os.path.isdir(name_or_path) else []
)
rel_path = COMFY_FILES[component]
found = resolve_comfy_file(
rel_path,
repo_id=FASTH3_REPO,
override_path=self.model_config.model_kwargs.get(f"{component}_path", None),
extra_roots=extra_roots,
status_fn=self.print_and_status_update,
local_only=True,
)
if found is not None:
return found
import huggingface_hub
target_dir = os.path.join(MODELS_PATH, "diffusion_models")
os.makedirs(target_dir, exist_ok=True)
self.print_and_status_update(
f"Downloading {os.path.basename(rel_path)} from {FASTH3_REPO}"
)
try:
return huggingface_hub.hf_hub_download(
repo_id=FASTH3_REPO,
filename=os.path.basename(rel_path),
local_dir=target_dir,
)
except Exception as e:
raise FileNotFoundError(
f"{os.path.basename(rel_path)} was not found locally or on "
f"{FASTH3_REPO}. Build it from the FastVideo release with:\n"
f" python scripts/convert_minimax_h2_fastvideo.py "
f"{FASTH3_SOURCE_REPO} {os.path.join(target_dir, os.path.basename(rel_path))}\n"
f"or point model_kwargs.dit_fasth3_path at an existing "
f"FastH3 int8-convrot file."
) from e
def _load_transformer(self) -> MiniMaxH3Transformer:
transformer = super()._load_transformer()
if not transformer.params.gate_compress:
raise ValueError(
"minimax_h3_vsa needs a VSA-trained checkpoint with "
"blocks.N.attn.to_gate_compress weights (FastH3); this file "
"has none. Use arch minimax_h3 for the base checkpoints."
)
transformer.vsa_sparsity = self.vsa_sparsity
return transformer
def _build_condition(
self, batch: "DataLoaderBatchDTO", latent_shape, device, dtype
):
# t2v only: no keyframe or reference conditioning
return None, None, (), ()
def get_base_model_version(self):
return "minimax_h3_vsa"
@property
def text_embedding_space_version(self):
# same Qwen3-VL presentation as the base model: share the embed cache
return "minimax_h3"

View File

@@ -546,12 +546,26 @@ def remap_sigma(
def build_sigma_schedule( def build_sigma_schedule(
num_inference_steps: int, shift: float = VIDEO_SIGMA_SHIFT num_inference_steps: int,
shift: float = VIDEO_SIGMA_SHIFT,
t1000_ladder: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
"""The released sampling grid: linspace(1, 0, steps + 1) through the """The released sampling grid: linspace(1, 0, steps + 1) through the
exponential shift, consecutive duplicates collapsed — `steps` yields exponential shift, consecutive duplicates collapsed — `steps` yields
`steps` model evaluations (the released repo counts the terminal 0 in `steps` model evaluations (the released repo counts the terminal 0 in
`steps`; we don't, so sample_steps means model evals).""" `steps`; we don't, so sample_steps means model evals).
base = torch.linspace(1.0, 0.0, num_inference_steps + 1, dtype=torch.float32)
``t1000_ladder`` uses FastH3's trained ladder instead: rounded indices on
the shared 1000-step grid ([999, 749, 500, 250] at 4 steps), one forward
per entry, each scheduler applying its own shift."""
if t1000_ladder:
base = (
torch.linspace(0.0, 999.0, num_inference_steps + 1, dtype=torch.float32)
.round()
.flip(0)
/ 1000.0
)
else:
base = torch.linspace(1.0, 0.0, num_inference_steps + 1, dtype=torch.float32)
sigmas = shift_sigma(base, shift) sigmas = shift_sigma(base, shift)
return torch.unique_consecutive(sigmas) return torch.unique_consecutive(sigmas)

View File

@@ -200,9 +200,11 @@ class MiniMaxH3Pipeline:
audio_rows = pack_audio_latents(audio_noise) # (1, 2*A, 32) audio_rows = pack_audio_latents(audio_noise) # (1, 2*A, 32)
# --- schedules ----------------------------------------------------- # --- schedules -----------------------------------------------------
sigmas_v = build_sigma_schedule(num_inference_steps, VIDEO_SIGMA_SHIFT).to( sigmas_v = build_sigma_schedule(
device num_inference_steps,
) VIDEO_SIGMA_SHIFT,
t1000_ladder=getattr(model, "t1000_sample_ladder", False),
).to(device)
# the audio schedule follows the video grid through the closed-form # the audio schedule follows the video grid through the closed-form
# shift remap so both streams sit at the same underlying position # shift remap so both streams sit at the same underlying position
sigmas_a = remap_sigma(sigmas_v, VIDEO_SIGMA_SHIFT, AUDIO_SIGMA_SHIFT) sigmas_a = remap_sigma(sigmas_v, VIDEO_SIGMA_SHIFT, AUDIO_SIGMA_SHIFT)
@@ -240,6 +242,8 @@ class MiniMaxH3Pipeline:
video_indices=video_indices, video_indices=video_indices,
audio_indices=audio_indices, audio_indices=audio_indices,
text_indices=text_indices, text_indices=text_indices,
# target-video token grid (patch 1x2x2); consumed only by VSA models
vsa_video_grid=(t_lat, h_lat // 2, w_lat // 2),
) )
v_video = video_pred[:, num_cond:].float() v_video = video_pred[:, num_cond:].float()
v_audio = audio_pred[:, layout.num_condition_audio_rows :].float() v_audio = audio_pred[:, layout.num_condition_audio_rows :].float()

View File

@@ -60,6 +60,9 @@ class MiniMaxH3TransformerParams:
norm_eps: float = 1e-5 norm_eps: float = 1e-5
qk_norm_eps: float = 1e-5 qk_norm_eps: float = 1e-5
final_norm_eps: float = 1e-5 final_norm_eps: float = 1e-5
# VSA-trained checkpoints (FastVideo FastH3) carry a per-token gate for
# the coarse compression branch: blocks.N.attn.to_gate_compress
gate_compress: bool = False
# "pruned" checkpoints (e.g. Comfy-Org *_pruned_*) replace the timestep # "pruned" checkpoints (e.g. Comfy-Org *_pruned_*) replace the timestep
# MLP with a small lookup table: ``adaln_t_table`` of shape # MLP with a small lookup table: ``adaln_t_table`` of shape
# (adaln_t_table_size, time_embed_dim) sampled by linear interpolation at # (adaln_t_table_size, time_embed_dim) sampled by linear interpolation at
@@ -146,7 +149,14 @@ class MiniMaxH3TimeEmbedder(nn.Module):
class MiniMaxH3Attention(nn.Module): class MiniMaxH3Attention(nn.Module):
"""Fused-QKV self-attention with per-head RMSNorm on q/k and partial RoPE.""" """Fused-QKV self-attention with per-head RMSNorm on q/k and partial RoPE."""
def __init__(self, hidden: int, heads: int, head_dim: int, qk_norm_eps: float): def __init__(
self,
hidden: int,
heads: int,
head_dim: int,
qk_norm_eps: float,
gate_compress: bool = False,
):
super().__init__() super().__init__()
self.heads = heads self.heads = heads
self.head_dim = head_dim self.head_dim = head_dim
@@ -155,12 +165,16 @@ class MiniMaxH3Attention(nn.Module):
self.q_norm = nn.RMSNorm(head_dim, eps=qk_norm_eps) self.q_norm = nn.RMSNorm(head_dim, eps=qk_norm_eps)
self.k_norm = nn.RMSNorm(head_dim, eps=qk_norm_eps) self.k_norm = nn.RMSNorm(head_dim, eps=qk_norm_eps)
self.out_proj = nn.Linear(inner, hidden, bias=False) self.out_proj = nn.Linear(inner, hidden, bias=False)
self.to_gate_compress = (
nn.Linear(hidden, inner, bias=False) if gate_compress else None
)
def forward( def forward(
self, self,
x: torch.Tensor, # (B, S, hidden) x: torch.Tensor, # (B, S, hidden)
rotary_emb=None, # (cos, sin) each (B, S, rot) or None rotary_emb=None, # (cos, sin) each (B, S, rot) or None
attn_mask: Optional[torch.Tensor] = None, # (B, 1, 1, S) bool, True = attend attn_mask: Optional[torch.Tensor] = None, # (B, 1, 1, S) bool, True = attend
vsa=None, # H3VSAContext or None (dense)
) -> torch.Tensor: ) -> torch.Tensor:
b, s, _ = x.shape b, s, _ = x.shape
q, k, v = self.qkv_proj(x).chunk(3, dim=-1) q, k, v = self.qkv_proj(x).chunk(3, dim=-1)
@@ -174,6 +188,13 @@ class MiniMaxH3Attention(nn.Module):
q = apply_rotary_emb(q, *rotary_emb) q = apply_rotary_emb(q, *rotary_emb)
k = apply_rotary_emb(k, *rotary_emb) k = apply_rotary_emb(k, *rotary_emb)
if vsa is not None and self.to_gate_compress is not None:
from .vsa import vsa_attention
gate = self.to_gate_compress(x).view(b, s, self.heads, self.head_dim)
out = vsa_attention(q, k, v, gate, vsa)
return self.out_proj(out.reshape(b, s, -1))
q = q.transpose(1, 2) q = q.transpose(1, 2)
k = k.transpose(1, 2) k = k.transpose(1, 2)
v = v.transpose(1, 2) v = v.transpose(1, 2)
@@ -282,7 +303,11 @@ class MiniMaxH3Block(nn.Module):
self.norm1 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps) self.norm1 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps)
self.norm2 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps) self.norm2 = nn.RMSNorm(p.hidden_size, eps=p.norm_eps)
self.attn = MiniMaxH3Attention( self.attn = MiniMaxH3Attention(
p.hidden_size, p.num_attention_heads, p.attention_head_dim, p.qk_norm_eps p.hidden_size,
p.num_attention_heads,
p.attention_head_dim,
p.qk_norm_eps,
gate_compress=p.gate_compress,
) )
self.mlp = MiniMaxH3Mlp(p.hidden_size, p.ffn_hidden_size) self.mlp = MiniMaxH3Mlp(p.hidden_size, p.ffn_hidden_size)
self.adaln_proj = MiniMaxH3AdalnProj( self.adaln_proj = MiniMaxH3AdalnProj(
@@ -301,6 +326,7 @@ class MiniMaxH3Block(nn.Module):
adaln_indices: torch.Tensor, # (B, S) long into the (M * 3) table adaln_indices: torch.Tensor, # (B, S) long into the (M * 3) table
rotary_emb, # (cos, sin) rotary_emb, # (cos, sin)
attn_mask: Optional[torch.Tensor] = None, attn_mask: Optional[torch.Tensor] = None,
vsa=None, # H3VSAContext or None (dense)
) -> torch.Tensor: ) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.adaln_proj(temb) self.adaln_proj(temb)
@@ -310,7 +336,9 @@ class MiniMaxH3Block(nn.Module):
h = self.norm1(x) * (1.0 + scale_msa[adaln_indices].to(dt)) + shift_msa[ h = self.norm1(x) * (1.0 + scale_msa[adaln_indices].to(dt)) + shift_msa[
adaln_indices adaln_indices
].to(dt) ].to(dt)
x = x + gate_msa[adaln_indices].to(dt) * self.attn(h, rotary_emb, attn_mask) x = x + gate_msa[adaln_indices].to(dt) * self.attn(
h, rotary_emb, attn_mask, vsa
)
h = self.norm2(x) * (1.0 + scale_mlp[adaln_indices].to(dt)) + shift_mlp[ h = self.norm2(x) * (1.0 + scale_mlp[adaln_indices].to(dt)) + shift_mlp[
adaln_indices adaln_indices
@@ -370,6 +398,7 @@ class MiniMaxH3Transformer(nn.Module, OstrisModelMixin):
# pruned checkpoint: factored timestep table instead of the MLP # pruned checkpoint: factored timestep table instead of the MLP
params.adaln_t_table_size = table.shape[0] params.adaln_t_table_size = table.shape[0]
params.time_embed_dim = table.shape[1] params.time_embed_dim = table.shape[1]
params.gate_compress = "blocks.0.attn.to_gate_compress.weight" in state_dict
return params return params
@classmethod @classmethod
@@ -416,6 +445,9 @@ class MiniMaxH3Transformer(nn.Module, OstrisModelMixin):
self.final_layer = MiniMaxH3FinalLayer(p) self.final_layer = MiniMaxH3FinalLayer(p)
self.gradient_checkpointing = False self.gradient_checkpointing = False
# None = dense attention. Set by the FastH3 model wrapper; only takes
# effect on gate_compress checkpoints when the caller passes the grid.
self.vsa_sparsity: Optional[float] = None
# float32 islands of the shipped checkpoint; used by the loader to keep # float32 islands of the shipped checkpoint; used by the loader to keep
# these keys at full precision when the rest is cast to bf16 # these keys at full precision when the rest is cast to bf16
@@ -469,6 +501,9 @@ class MiniMaxH3Transformer(nn.Module, OstrisModelMixin):
video_indices: torch.Tensor, # (Nv,) long positions of video rows in the pack video_indices: torch.Tensor, # (Nv,) long positions of video rows in the pack
audio_indices: torch.Tensor, # (Na,) long audio_indices: torch.Tensor, # (Na,) long
text_indices: torch.Tensor, # (L,) long text_indices: torch.Tensor, # (L,) long
vsa_video_grid: Optional[
Tuple[int, int, int]
] = None, # target-video token grid (t, h, w)
): ):
"""Returns (video_out (B, Nv, 96), audio_out (B, Na, 32)) — the """Returns (video_out (B, Nv, 96), audio_out (B, Na, 32)) — the
data-ward velocity ``clean - noise`` for every row, in input order. data-ward velocity ``clean - noise`` for every row, in input order.
@@ -510,6 +545,28 @@ class MiniMaxH3Transformer(nn.Module, OstrisModelMixin):
temb = self._time_embedding(unique_t) temb = self._time_embedding(unique_t)
adaln_indices = inverse * MODALITY_NUM + token_tags.clamp(min=0) adaln_indices = inverse * MODALITY_NUM + token_tags.clamp(min=0)
vsa_ctx = None
if (
self.vsa_sparsity is not None
and self.params.gate_compress
and vsa_video_grid is not None
):
from .vsa import build_vsa_context, vsa_is_available
# no triton -> warn once and run every block dense instead
vsa_ctx = (
None
if not vsa_is_available()
else build_vsa_context(
seq_len=seq_len,
num_text_rows=int(text_indices.shape[0]),
video_grid=tuple(int(g) for g in vsa_video_grid),
token_tags=token_tags,
sparsity=self.vsa_sparsity,
device=x.device,
)
)
for block in self.blocks: for block in self.blocks:
if torch.is_grad_enabled() and self.gradient_checkpointing: if torch.is_grad_enabled() and self.gradient_checkpointing:
x = checkpoint( x = checkpoint(
@@ -519,10 +576,11 @@ class MiniMaxH3Transformer(nn.Module, OstrisModelMixin):
adaln_indices, adaln_indices,
rotary_emb, rotary_emb,
attn_mask, attn_mask,
vsa_ctx,
use_reentrant=False, use_reentrant=False,
) )
else: else:
x = block(x, temb, adaln_indices, rotary_emb, attn_mask) x = block(x, temb, adaln_indices, rotary_emb, attn_mask, vsa_ctx)
video_all, audio_all = self.final_layer(x, temb, inverse) video_all, audio_all = self.final_layer(x, temb, inverse)
video_out = video_all.index_select(1, video_indices) video_out = video_all.index_select(1, video_indices)

View File

@@ -0,0 +1,343 @@
"""FastVideo VSA (Video Sparse Attention) for MiniMax-H3, on FlexAttention.
Replicates ``fastvideo/attention/backends/video_sparse_attn_h3.py`` without
the fastvideo_kernel CUDA/Triton extensions. The packed
``[text | condition | audio | video]`` sequence is re-tiled so every 64
contiguous rows form one tile: prefix segments chunk in order (tiles never
straddle a segment boundary) and the video tail becomes (4, 4, 4) 3D tiles,
zero-padding ragged edges. Per-tile fp32 mean pooling of q/k scores every
tile pair; each video query tile keeps the top
``ceil((1 - sparsity) * num_video_tiles)`` video tiles ("exempt" mode:
prefix tiles are always visible and prefix queries run dense). The fine
stage is exact attention over the selected tiles via ``flex_attention`` with
a 64-token BlockMask; the compression branch adds
``gate * softmax(scores) @ v_pooled`` per tile, gate from
``to_gate_compress``.
The FastH3 4-step checkpoints were distilled WITH this policy (sparsity 0.9,
tile 64 is the trained geometry), so training and sampling both run it, at
every sequence length.
"""
import math
import os
from dataclasses import dataclass
from functools import lru_cache
from typing import Optional, Tuple
import torch
from torch.nn.attention.flex_attention import BlockMask, flex_attention
TILE_SHAPE = (4, 4, 4)
TILE = 64
# fine-stage backend: "auto" runs FastVideo's vendored Triton block-sparse
# kernels when eligible (uniform batch, fp16/bf16, default scale) and
# FlexAttention otherwise; "flex" / "fv_triton" force one side.
VSA_BACKEND = os.environ.get("AITK_H3_VSA_BACKEND", "auto")
_flex_compiled = None
_triton_ok: Optional[bool] = None
def vsa_is_available() -> bool:
"""Compiled FlexAttention needs triton (inductor GPU codegen). Without it
the only flex path is the eager fallback, which materializes the full
score matrix — unusable at video lengths — so VSA is disabled instead."""
global _triton_ok
if _triton_ok is None:
try:
import triton # noqa: F401
_triton_ok = True
except Exception:
_triton_ok = False
print(
"WARNING: triton is not installed; VSA sparse attention is "
"disabled and MiniMax-H3 FastH3 falls back to dense attention "
"(FastVideo's dense mode, compression gate off). Install "
"triton to run the checkpoint's trained sparse policy."
)
return _triton_ok
def _get_flex(disable_compile: bool = False):
if disable_compile:
return flex_attention
global _flex_compiled
if _flex_compiled is None:
_flex_compiled = torch.compile(flex_attention, dynamic=False)
return _flex_compiled
@dataclass(frozen=True)
class H3VSAGeometry:
seq_len: int
n_tiles: int
num_prefix_tiles: int
num_video_tiles: int
tile_sizes: torch.Tensor # (n_tiles,) long, live rows per tile
untile_index: torch.Tensor # (seq_len,) long, packed row -> padded slot
slot_valid: torch.Tensor # (n_tiles * TILE,) bool
@dataclass
class H3VSAContext:
geometry: H3VSAGeometry
sparsity: float
token_valid: Optional[torch.Tensor] # (B, S) bool, None = all live
def _video_tile_partition(
grid: Tuple[int, int, int], device
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Row indices of the (T, H, W) token grid grouped into (4, 4, 4) tiles in
(t-tile, h-tile, w-tile) raster order, plus each tile's live-row count."""
t, h, w = grid
ts, hs, ws = TILE_SHAPE
indices = torch.arange(t * h * w, device=device, dtype=torch.long).reshape(t, h, w)
parts, sizes = [], []
for ti in range(math.ceil(t / ts)):
for hi in range(math.ceil(h / hs)):
for wi in range(math.ceil(w / ws)):
block = indices[
ti * ts : min(ti * ts + ts, t),
hi * hs : min(hi * hs + hs, h),
wi * ws : min(wi * ws + ws, w),
].flatten()
parts.append(block)
sizes.append(block.numel())
return torch.cat(parts), torch.tensor(sizes, dtype=torch.long, device=device)
@lru_cache(maxsize=16)
def _geometry_cached(
prefix_segments: Tuple[int, ...], video_grid: Tuple[int, int, int], device_str: str
) -> H3VSAGeometry:
device = torch.device(device_str)
prefix_len = sum(prefix_segments)
prefix_sizes = []
for segment in prefix_segments:
full, rem = divmod(segment, TILE)
prefix_sizes.extend([TILE] * full)
if rem:
prefix_sizes.append(rem)
num_prefix_tiles = len(prefix_sizes)
video_partition, video_sizes = _video_tile_partition(video_grid, device)
num_video_tiles = int(video_sizes.numel())
partition = torch.cat(
[
torch.arange(prefix_len, device=device, dtype=torch.long),
video_partition + prefix_len,
]
)
tile_sizes = torch.cat(
[torch.tensor(prefix_sizes, dtype=torch.long, device=device), video_sizes]
)
n_tiles = int(tile_sizes.numel())
seq_len = int(partition.numel())
# padded slot of the k-th row of the tile-ordered sequence: tiles occupy
# TILE slots each, live rows sit at the front of their tile
starts = torch.arange(n_tiles, device=device, dtype=torch.long) * TILE
shift = torch.cat([tile_sizes.new_zeros(1), tile_sizes.cumsum(0)[:-1]])
intra = torch.arange(seq_len, device=device) - shift.repeat_interleave(tile_sizes)
non_pad = starts.repeat_interleave(tile_sizes) + intra
untile_index = non_pad[torch.argsort(partition)]
slot_valid = torch.zeros(n_tiles * TILE, dtype=torch.bool, device=device)
slot_valid[non_pad] = True
return H3VSAGeometry(
seq_len=seq_len,
n_tiles=n_tiles,
num_prefix_tiles=num_prefix_tiles,
num_video_tiles=num_video_tiles,
tile_sizes=tile_sizes,
untile_index=untile_index,
slot_valid=slot_valid,
)
def build_vsa_context(
seq_len: int,
num_text_rows: int,
video_grid: Tuple[int, int, int],
token_tags: torch.Tensor, # (B, S) long, -1 marks pad rows
sparsity: float,
device: torch.device,
) -> H3VSAContext:
"""Context for one forward. The video tokens are the sequence tail; the
prefix splits as (text, everything between text and video) — for the t2v
packing this is FastVideo's (text, audio) segmentation exactly."""
video_rows = video_grid[0] * video_grid[1] * video_grid[2]
prefix_rest = seq_len - num_text_rows - video_rows
if prefix_rest < 0:
raise ValueError(
f"VSA video grid {video_grid} ({video_rows} rows) does not fit the "
f"packed sequence ({seq_len} rows, {num_text_rows} text)"
)
prefix_segments = tuple(s for s in (num_text_rows, prefix_rest) if s > 0)
geometry = _geometry_cached(prefix_segments, tuple(video_grid), str(device))
token_valid = None
is_pad = token_tags < 0
if bool(is_pad.any()):
token_valid = ~is_pad
return H3VSAContext(geometry=geometry, sparsity=sparsity, token_valid=token_valid)
def compute_topk(sparsity: float, num_blocks: int) -> int:
"""Video tiles each video query tile keeps, clamped to [1, num_blocks]."""
return max(1, min(math.ceil((1.0 - sparsity) * num_blocks), num_blocks))
def vsa_attention(
q: torch.Tensor, # (B, S, H, D), post qk-norm and rope
k: torch.Tensor,
v: torch.Tensor,
gate: Optional[torch.Tensor], # (B, S, H, D) from to_gate_compress, or None
ctx: H3VSAContext,
scale: Optional[float] = None,
disable_compile: bool = False,
backend: Optional[str] = None, # None -> VSA_BACKEND
) -> torch.Tensor:
b, s, heads, d = q.shape
g = ctx.geometry
if s != g.seq_len:
raise ValueError(f"VSA geometry was built for {g.seq_len} rows, got {s}")
if scale is None:
scale = d**-0.5
idx = g.untile_index
n = g.n_tiles
s_pad = n * TILE
token_valid = ctx.token_valid
if token_valid is not None:
# dead rows must not enter pooling or act as keys
live = token_valid[..., None, None].to(q.dtype)
q, k, v = q * live, k * live, v * live
slot_valid = torch.zeros(b, s_pad, dtype=torch.bool, device=q.device)
slot_valid[:, idx] = token_valid
sizes = slot_valid.view(b, n, TILE).sum(-1).to(torch.float32) # (B, n)
else:
slot_valid = g.slot_valid.unsqueeze(0).expand(b, -1)
sizes = g.tile_sizes.to(torch.float32).unsqueeze(0) # (1, n)
def tile_rows(x):
buf = x.new_zeros(b, s_pad, heads, d)
buf[:, idx] = x
return buf
qt, kt, vt = tile_rows(q), tile_rows(k), tile_rows(v)
def pool(x):
# pad slots are zero, so a plain fp32 sum / live count is a masked mean
p = x.view(b, n, TILE, heads, d).sum(2, dtype=torch.float32)
return (p / sizes.clamp(min=1.0)[..., None, None]).permute(0, 2, 1, 3)
scores = torch.matmul(pool(qt), pool(kt).transpose(-2, -1)) * scale # (B,H,n,n)
npre = g.num_prefix_tiles
k_vid = compute_topk(ctx.sparsity, g.num_video_tiles)
if k_vid == g.num_video_tiles:
mask = torch.ones(b, heads, n, n, dtype=torch.bool, device=q.device)
else:
mask = torch.zeros(b, heads, n, n, dtype=torch.bool, device=q.device)
top = scores[..., npre:].topk(k_vid, dim=-1).indices + npre
mask.scatter_(-1, top, True)
mask[..., :npre] = True # prefix tiles always visible ("exempt" mode)
mask[:, :, :npre, :] = True # prefix queries run dense
def finish(out):
"""Shared tail: gated compression branch, then un-tile to packed order."""
if gate is not None:
# compression branch: dense attention over the pooled tiles,
# broadcast to each tile's rows, scaled by the learned gate
gt = tile_rows(gate)
dead = sizes <= 0 # a tile of only pad rows must not be attended
cs = scores
if bool(dead.any()):
cs = cs.masked_fill(dead.view(-1, 1, 1, n), torch.finfo(cs.dtype).min)
oc = torch.matmul(torch.softmax(cs.float(), dim=-1), pool(vt))
oc = oc.permute(0, 2, 1, 3).to(out.dtype) # (B, n, H, D)
out = (
out.view(b, n, TILE, heads, d)
+ oc.unsqueeze(2) * gt.view(b, n, TILE, heads, d)
).view(b, s_pad, heads, d)
return out[:, idx]
if backend is None:
backend = VSA_BACKEND
use_fv = backend != "flex" and (
token_valid is None # per-item pad rows can't be expressed as tile sizes
and q.is_cuda
and q.dtype in (torch.bfloat16, torch.float16)
and scale == d**-0.5 # the kernel hardcodes 1/sqrt(D)
)
if backend == "fv_triton" and not use_fv:
raise ValueError(
"fv_triton backend needs a uniform (pad-free) batch, fp16/bf16 "
"CUDA tensors and the default scale"
)
if use_fv:
from . import vsa_kernels
out = vsa_kernels.block_sparse_attn(
qt.permute(0, 2, 1, 3),
kt.permute(0, 2, 1, 3),
vt.permute(0, 2, 1, 3),
mask,
g.tile_sizes,
).permute(0, 2, 1, 3) # (B, S_pad, H, D)
return finish(out)
# fully-live selected tiles skip the mask_mod; ragged/padded tiles keep it
col_full = (sizes >= TILE).view(-1, 1, 1, n)
m_full = mask & col_full
m_part = mask & ~col_full
def to_blocks(m):
num = m.sum(-1, dtype=torch.int32)
order = m.to(torch.uint8).argsort(dim=-1, descending=True, stable=True)
return num, order.to(torch.int32)
valid = slot_valid
tile_mask = mask
def mask_mod(bi, hi, qi, ki):
# the FULL mask truth: eager flex ignores the kv block lists and
# evaluates only mask_mod, so it must carry tile selection too
return tile_mask[bi, hi, qi // TILE, ki // TILE] & valid[bi, ki]
part_num, part_idx = to_blocks(m_part)
full_num, full_idx = to_blocks(m_full)
block_mask = BlockMask.from_kv_blocks(
part_num,
part_idx,
full_kv_num_blocks=full_num,
full_kv_indices=full_idx,
BLOCK_SIZE=TILE,
mask_mod=mask_mod,
)
out = _get_flex(disable_compile)(
qt.permute(0, 2, 1, 3),
kt.permute(0, 2, 1, 3),
vt.permute(0, 2, 1, 3),
block_mask=block_mask,
scale=scale,
# inductor's flex kernels need their tile sizes to divide the 64-token
# mask blocks; the defaults (128) reject BLOCK_SIZE=64
kernel_options={
"BLOCK_M": 64,
"BLOCK_N": 64,
"BLOCK_M1": 32,
"BLOCK_N1": 64,
"BLOCK_M2": 64,
"BLOCK_N2": 32,
},
).permute(0, 2, 1, 3) # (B, S_pad, H, D)
return finish(out)

View File

@@ -0,0 +1,144 @@
"""Autograd glue for the vendored FastVideo Triton block-sparse kernels.
Adapted from FastVideo's ``fastvideo_kernel/block_sparse_attn.py`` (Apache-2.0)
with the compiled sm90/sm100a backends stripped: only the pure-Triton fwd+bwd
pair remains, JIT-compiled on first call. Custom ops live in their own
``aitk_h3_vsa::`` namespace so a real fastvideo-kernel install cannot clash.
Public entry: ``block_sparse_attn(q, k, v, block_map, variable_block_sizes)``
with q/k/v ``[B, H, S_pad, D]`` (S_pad a multiple of 64), block_map a bool
``[B, H, n_tiles, n_tiles]``, variable_block_sizes int32 ``[n_tiles]`` live
rows per tile (live rows at the front of each 64-token tile).
"""
from typing import Tuple
import torch
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
if not t.is_cuda:
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
if t.dtype != torch.int32:
t = t.to(torch.int32)
if not t.is_contiguous():
t = t.contiguous()
return t
@torch.library.custom_op(
"aitk_h3_vsa::block_sparse_attn_triton",
mutates_args=(),
device_types="cuda",
)
def _block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from .block_sparse_attn_triton import triton_block_sparse_attn_forward
o, M = triton_block_sparse_attn_forward(
q.contiguous(),
k.contiguous(),
v.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o, M
@torch.library.register_fake("aitk_h3_vsa::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(q, k, v, q2k_idx, q2k_num, variable_block_sizes):
o = torch.empty_like(q)
M = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32
)
return o, M
@torch.library.custom_op(
"aitk_h3_vsa::block_sparse_attn_backward_triton",
mutates_args=(),
device_types="cuda",
)
def _block_sparse_attn_backward_triton(
grad_output: torch.Tensor,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
from .block_sparse_attn_triton import triton_block_sparse_attn_backward
from .index import invert_indices
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
# o/M are kernel outputs so are already contiguous.
dq, dk, dv = triton_block_sparse_attn_backward(
grad_output.contiguous(),
q.contiguous(),
k.contiguous(),
v.contiguous(),
o,
M,
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
variable_block_sizes,
)
return dq, dk, dv
@torch.library.register_fake("aitk_h3_vsa::block_sparse_attn_backward_triton")
def _block_sparse_attn_backward_triton_fake(
grad_output, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
):
return torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
def _setup_context(ctx, inputs, output):
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
def _backward(ctx, grad_o, grad_M):
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = _block_sparse_attn_backward_triton(
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
_block_sparse_attn_triton.register_autograd(_backward, setup_context=_setup_context)
def block_sparse_attn(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
"""Block-sparse attention with autograd from a bool tile map. Returns o."""
from .index import map_to_index
q2k_idx, q2k_num = map_to_index(block_map)
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
variable_block_sizes = _as_int32_contig(
variable_block_sizes, "variable_block_sizes"
)
o, _ = _block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
return o

View File

@@ -0,0 +1,893 @@
"""Vendored from FastVideo (https://github.com/hao-ai-lab/FastVideo),
``fastvideo-kernel/python/fastvideo_kernel/triton_kernels/block_sparse_attn_triton.py``.
Apache-2.0; Copyright the FastVideo team. Vendored unmodified so the
MiniMax-H3 VSA fine stage can run FastVideo's trained-policy Triton
block-sparse kernels (fwd + bwd) without the fastvideo-kernel package."""
"""
Fused Attention
===============
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
(https://tridao.me/publications/flash2/flash2.pdf)
Credits: OpenAI kernel team
"""
import torch
import triton
import triton.language as tl
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
import math # small utility needed by the sparse wrapper
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
# BLOCK_M / BLOCK_N are fixed at 64 because they are structural, not tunable:
# the kernel indexes the top-k list per BLOCK_M q-tile and addresses keys as
# kv_idx * BLOCK_N, so both must match the granularity q2k_index and
# variable_block_sizes were built at.
#
# num_stages / num_warps ARE free, and the previous {3, 4, 7} was inherited from
# the upstream tutorial rather than tuned here. It skips 5 and 6; on Blackwell
# (sm_121) the optimum is num_stages=5, so the search could not reach it. Both
# block paths independently select 5 once it is available. Autotune still picks
# per architecture, so other GPUs re-tune rather than inheriting this choice.
configs = [
triton.Config({"BLOCK_M": BM, "BLOCK_N": BN}, num_stages=s, num_warps=w)
for BM in [64]
for BN in [64]
for s in [2, 3, 4, 5, 6, 7]
for w in [4, 8]
]
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX_Q", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(
Q,
K,
V,
sm_scale, #
q2k_index,
q2k_num,
max_kv_blks, #
variable_block_sizes,
M,
Out, #
stride_qz,
stride_qh,
stride_qm,
stride_qk,
stride_kz,
stride_kh,
stride_kn,
stride_kk,
stride_vz,
stride_vh,
stride_vk,
stride_vn,
stride_oz,
stride_oh,
stride_om,
stride_on,
Z,
H,
N_CTX_Q, #
N_CTX_KV, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
STAGE: tl.constexpr,
):
"""
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
(32×64 and 64×32) – memory footprint unchanged.
"""
# ----- program-id mapping -----
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX_Q // BLOCK_M
meta_base = (b * H + h) * q_tiles + q_blk
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
# Note: when q and kv have different sequence lengths, their per-(batch,head)
# strides differ, so we must compute separate base offsets.
q_off = b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh
k_off = b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh
v_off = b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh
o_off = b.to(tl.int64) * stride_oz + h.to(tl.int64) * stride_oh
Q_ptr = tl.make_block_ptr(
base=Q + q_off,
shape=(N_CTX_Q, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
K_base = tl.make_block_ptr(
base=K + k_off,
shape=(HEAD_DIM, N_CTX_KV),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N),
order=(0, 1),
)
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
V_base = tl.make_block_ptr(
base=V + v_off,
shape=(N_CTX_KV, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=v_order,
)
O_ptr = tl.make_block_ptr(
base=Out + o_off,
shape=(N_CTX_Q, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
# ----- accumulators -----
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
qk_scale = sm_scale * 1.44269504 # 1/ln2
q = tl.load(Q_ptr)
# ----- sparse loop over valid K/V tiles -----
for i in range(0, kv_blocks):
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
block_size = tl.load(variable_block_sizes + kv_idx)
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
k = tl.load(K_ptr)
qk = tl.dot(q, k)
# mask out invalid columns
mask = tl.arange(0, BLOCK_N) < block_size
qk = tl.where(mask[None, :], qk, -float("inf"))
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
l_ij = tl.sum(p, 1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
v = tl.load(V_ptr)
acc = tl.dot(p.to(tl.bfloat16), v, acc)
m_i = m_ij
# ----- epilogue -----
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX_Q + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
@triton.jit
def _attn_bwd_preprocess(
O,
DO, #
Delta, #
Z,
H,
N_CTX, #
BLOCK_M: tl.constexpr,
HEAD_DIM: tl.constexpr, #
):
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
off_hz = tl.program_id(1)
off_n = tl.arange(0, HEAD_DIM)
# load
o = tl.load(
O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]
)
do = tl.load(
DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]
).to(tl.float32)
delta = tl.sum(o * do, axis=1)
# write-back
tl.store(Delta + off_hz * N_CTX + off_m, delta)
# The main inner-loop logic for computing dK and dV.
@triton.jit
def _attn_bwd_dkdv(
dk,
dv, #
Q,
k,
v,
sm_scale, #
DO, #
M,
D, #
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_tok,
stride_d, #
H,
N_CTX_KV,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
# Filled in by the wrapper.
start_n,
start_m,
num_steps,
):
offs_m = start_m + tl.arange(0, BLOCK_M1)
offs_n = start_n + tl.arange(0, BLOCK_N1)
offs_k = tl.arange(0, HEAD_DIM)
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
step_m = BLOCK_M1
kv_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
kv_tiles = N_CTX_KV // BLOCK_N1
meta_base = (b * H + h) * kv_tiles + kv_blk
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
block_size = tl.load(variable_block_sizes + kv_blk)
for blk_idx in range(q_blocks * 2):
block_sparse_offset = (
tl.load(q_ptr + blk_idx // 2).to(tl.int32) * 2 + blk_idx % 2
) * step_m
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
# Recompute logits exactly as the forward does: raw bf16 operands into
# the dot, fp32 scale after accumulation. A bf16 pre-scaled K perturbs
# the recomputed logits relative to the saved M by an error
# proportional to |logit|, which exp2 amplifies into arbitrarily wrong
# probabilities at large activations.
qkT = tl.dot(k, qT) * (sm_scale * 1.4426950408889634)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
# Compute dV.
ppT = pT
ppT = ppT.to(tl.bfloat16)
dv += tl.dot(ppT, do)
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# Compute dP and dS.
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
dsT = pT * (dpT - Di[None, :])
dsT = dsT.to(tl.bfloat16)
dk += tl.dot(dsT, tl.trans(qT))
# Increment pointers.
return dk, dv
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq(
dq,
q,
K,
V, #
do,
m,
D,
sm_scale,
# shared by Q/K/V/DO.
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
# Filled in by the wrapper.
start_m,
start_n,
num_steps,
):
offs_m = start_m + tl.arange(0, BLOCK_M2)
offs_n = start_n + tl.arange(0, BLOCK_N2)
offs_k = tl.arange(0, HEAD_DIM)
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
step_n = BLOCK_N2
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M2
meta_base = (b * H + h) * q_tiles + q_blk
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
for blk_idx in range(kv_blocks * 2):
kv_idx = tl.load(kv_ptr + blk_idx // 2).to(tl.int32)
# variable_block_sizes is defined per KV block (tile). Mask must therefore
# use kv_idx (not q_blk). Also, because we split each 64-token block into
# two 32-token halves, the mask must account for the half-block offset.
block_size = tl.load(variable_block_sizes + kv_idx).to(tl.int32)
half = (blk_idx % 2).to(tl.int32)
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
p = tl.math.exp2(qk - m)
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
mask = offs_in_block < block_size
p = tl.where(mask[None, :], p, 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@triton.jit
def _attn_bwd(
Q,
K,
V,
sm_scale, #
DO, #
DQ,
DK,
DV, #
M,
D,
q2k_index,
q2k_num,
max_kv_blks,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_z,
stride_h,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
):
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
off_chz = (bhid * N_CTX).to(tl.int64)
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
pid = tl.program_id(0)
# offset pointers for batch/head
Q += adj
K += adj
V += adj
DO += adj
DQ += adj
DK += adj
DV += adj
M += off_chz
D += off_chz
# load scales
offs_k = tl.arange(0, HEAD_DIM)
start_n = pid * BLOCK_N1
start_m = 0
offs_n = start_n + tl.arange(0, BLOCK_N1)
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
num_steps = N_CTX // BLOCK_M1
dk, dv = _attn_bwd_dkdv( #
dk,
dv, #
Q,
k,
v,
sm_scale, #
DO, #
M,
D, #
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M1,
BLOCK_N1,
HEAD_DIM, #
start_n,
start_m,
num_steps, #
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv)
# Write back dK.
dk *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk)
# THIS BLOCK DOES DQ:
start_m = pid * BLOCK_M2
end_n = 0
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)
m = m[:, None]
num_steps = N_CTX // BLOCK_N2
dq = _attn_bwd_dq(
dq,
q,
K,
V, #
do,
m,
D, #
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M2,
BLOCK_N2,
HEAD_DIM, #
start_m,
end_n,
num_steps, #
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= sm_scale
tl.store(dq_ptrs, dq)
@triton.jit
def _attn_bwd_dkdv_kernel(
Q,
K,
V,
sm_scale, #
DO, #
DK,
DV, #
M,
D,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared token/dim strides (assumed contiguous along token and dim)
stride_tok,
stride_d, #
# batch/head strides (may differ between Q and KV)
stride_qz,
stride_qh,
stride_kz,
stride_kh,
stride_vz,
stride_vh,
stride_doz,
stride_doh,
stride_dkz,
stride_dkh,
stride_dvz,
stride_dvh,
H,
N_CTX_Q,
N_CTX_KV,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr,
):
"""
Backward kernel that computes dK and dV for each KV block (64 tokens).
Grid:
pid0: kv_blk in [0, N_CTX_KV/BLOCK_N1)
pid2: fused (batch, head) in [0, B*H)
"""
bhid = tl.program_id(2)
b = bhid // H
h = bhid % H
kv_blk = tl.program_id(0)
q_adj = b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh
kv_adj_k = b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh
kv_adj_v = b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh
do_adj = b.to(tl.int64) * stride_doz + h.to(tl.int64) * stride_doh
dk_adj = b.to(tl.int64) * stride_dkz + h.to(tl.int64) * stride_dkh
dv_adj = b.to(tl.int64) * stride_dvz + h.to(tl.int64) * stride_dvh
Q = Q + q_adj
K = K + kv_adj_k
V = V + kv_adj_v
DO = DO + do_adj
DK = DK + dk_adj
DV = DV + dv_adj
# M and D (delta) are always sized by Q length.
M = M + (bhid * N_CTX_Q).to(tl.int64)
D = D + (bhid * N_CTX_Q).to(tl.int64)
offs_k = tl.arange(0, HEAD_DIM)
start_n = kv_blk * BLOCK_N1
offs_n = start_n + tl.arange(0, BLOCK_N1)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
dv_acc = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk_acc = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
num_steps = N_CTX_Q // BLOCK_M1
dk_acc, dv_acc = _attn_bwd_dkdv(
dk_acc,
dv_acc,
Q,
k,
v,
sm_scale,
DO,
M,
D,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
stride_tok,
stride_d,
H,
N_CTX_KV,
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1,
HEAD_DIM=HEAD_DIM,
start_n=start_n,
start_m=0,
num_steps=num_steps,
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv_acc)
dk_acc *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk_acc)
@triton.jit
def _attn_bwd_dq_kernel(
Q,
K,
V,
sm_scale,
DO, #
DQ,
M,
D,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
# shared token/dim strides (assumed contiguous along token and dim)
stride_tok,
stride_d, #
# batch/head strides (may differ between Q and KV)
stride_qz,
stride_qh,
stride_kz,
stride_kh,
stride_vz,
stride_vh,
stride_doz,
stride_doh,
stride_dqz,
stride_dqh,
H,
N_CTX_Q,
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
):
"""
Backward kernel that computes dQ for each Q block (64 tokens).
Grid:
pid0: q_blk in [0, N_CTX_Q/BLOCK_M2)
pid2: fused (batch, head) in [0, B*H)
"""
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
b = bhid // H
h = bhid % H
q_blk = tl.program_id(0)
q_adj = b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh
kv_adj_k = b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh
kv_adj_v = b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh
do_adj = b.to(tl.int64) * stride_doz + h.to(tl.int64) * stride_doh
dq_adj = b.to(tl.int64) * stride_dqz + h.to(tl.int64) * stride_dqh
Q = Q + q_adj
K = K + kv_adj_k
V = V + kv_adj_v
DO = DO + do_adj
DQ = DQ + dq_adj
M = M + (bhid * N_CTX_Q).to(tl.int64)
D = D + (bhid * N_CTX_Q).to(tl.int64)
offs_k = tl.arange(0, HEAD_DIM)
start_m = q_blk * BLOCK_M2
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)[:, None]
dq_acc = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
num_steps = 0 # unused in _attn_bwd_dq
dq_acc = _attn_bwd_dq(
dq_acc,
q,
K,
V,
do,
m,
D,
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d,
H,
N_CTX_Q,
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2,
HEAD_DIM=HEAD_DIM,
start_m=start_m,
start_n=0,
num_steps=num_steps,
)
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq_acc *= sm_scale
tl.store(dq_ptrs, dq_acc)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
B, H, Tq, D = q.shape
Tkv = k.shape[2]
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert Tq % 64 == 0, f"q length must be a multiple of 64, but got {Tq}"
assert Tkv % 64 == 0, f"kv length must be a multiple of 64, but got {Tkv}"
assert q2k_num.shape[-1] == Tq // 64, (
f"shape mismatch, Tq // 64 = {Tq // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
)
assert variable_block_sizes.numel() == Tkv // 64, (
f"shape mismatch, variable_block_sizes must have length {Tkv // 64}, "
f"got {variable_block_sizes.numel()}"
)
o = torch.empty_like(q)
M = torch.empty((B, H, Tq), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(Tq, 64), B * H, 1)
_attn_fwd_sparse[grid](
q,
k,
v,
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
M,
o,
q.stride(0),
q.stride(1),
q.stride(2),
q.stride(3),
k.stride(0),
k.stride(1),
k.stride(2),
k.stride(3),
v.stride(0),
v.stride(1),
v.stride(2),
v.stride(3),
o.stride(0),
o.stride(1),
o.stride(2),
o.stride(3),
B,
H,
Tq,
Tkv,
HEAD_DIM=D,
STAGE=3,
)
return o, M
def triton_block_sparse_attn_backward(
do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes
):
assert do.is_contiguous()
B, H, Tq, D = q.shape
Tkv = k.shape[2]
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
BATCH, N_HEAD = q.shape[:2]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
# K stays raw: the backward kernels apply sm_scale in fp32 after the dot,
# matching the forward's rounding exactly. (A bf16 pre-scaled K perturbs
# the recomputed logits vs the saved M; exp2 turns that into unboundedly
# wrong probabilities at large activations.)
arg_k = k
PRE_BLOCK = 64
assert Tq % PRE_BLOCK == 0
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o,
do, #
delta, #
BATCH,
N_HEAD,
Tq, #
BLOCK_M=PRE_BLOCK,
HEAD_DIM=D, #
)
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
# dK/dV kernel: grid over KV blocks
grid_kv = (Tkv // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd_dkdv_kernel[grid_kv](
q,
arg_k,
v,
sm_scale,
do,
dk,
dv,
M,
delta,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
q.stride(2),
q.stride(3),
q.stride(0),
q.stride(1),
arg_k.stride(0),
arg_k.stride(1),
v.stride(0),
v.stride(1),
do.stride(0),
do.stride(1),
dk.stride(0),
dk.stride(1),
dv.stride(0),
dv.stride(1),
N_HEAD,
Tq,
Tkv,
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1,
HEAD_DIM=D,
)
# dQ kernel: grid over Q blocks
grid_q = (Tq // BLOCK_M2, 1, BATCH * N_HEAD)
_attn_bwd_dq_kernel[grid_q](
q,
arg_k,
v,
sm_scale,
do,
dq,
M,
delta,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
q.stride(2),
q.stride(3),
q.stride(0),
q.stride(1),
arg_k.stride(0),
arg_k.stride(1),
v.stride(0),
v.stride(1),
do.stride(0),
do.stride(1),
dq.stride(0),
dq.stride(1),
N_HEAD,
Tq,
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2,
HEAD_DIM=D,
)
return dq, dk, dv

View File

@@ -0,0 +1,288 @@
"""Vendored from FastVideo (https://github.com/hao-ai-lab/FastVideo),
``fastvideo-kernel/python/fastvideo_kernel/triton_kernels/index.py``.
Apache-2.0; Copyright the FastVideo team. Vendored unmodified so the
MiniMax-H3 VSA fine stage can run FastVideo's trained-policy Triton
block-sparse kernels (fwd + bwd) without the fastvideo-kernel package."""
## pytorch sdpa version of block sparse ##
from typing import Tuple
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = (
index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
)
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = (
index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
)
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr
+ b * index_num_bs_stride
+ h * index_num_h_stride
+ q * index_num_q_stride,
num,
)
def topk_index_to_map(
index: torch.Tensor, num_kv_blocks: int, transpose_map: bool = False
):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros(
(bs, h, num_q_blocks, num_kv_blocks), dtype=torch.bool, device=index.device
)
else:
block_map = torch.zeros(
(bs, h, num_kv_blocks, num_q_blocks), dtype=torch.bool, device=index.device
)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full(
(block_map.shape), -1, dtype=torch.int32, device=block_map.device
)
index_num = torch.empty(
(bs, h, num_q_blocks), dtype=torch.int32, device=block_map.device
)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
@triton.jit
def _invert_indices_kernel(
q2k_idx_ptr,
q2k_num_ptr,
k2q_idx_ptr,
k2q_num_ptr,
q2k_idx_b,
q2k_idx_h,
q2k_idx_q,
q2k_idx_k,
q2k_num_b,
q2k_num_h,
q2k_num_q,
k2q_idx_b,
k2q_idx_h,
k2q_idx_k,
k2q_idx_q,
k2q_num_b,
k2q_num_h,
k2q_num_k,
MAX_KV_PER_Q: tl.constexpr,
):
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
pid_q = tl.program_id(2)
n = tl.load(q2k_num_ptr + pid_b * q2k_num_b + pid_h * q2k_num_h + pid_q * q2k_num_q)
q2k_row = q2k_idx_ptr + pid_b * q2k_idx_b + pid_h * q2k_idx_h + pid_q * q2k_idx_q
for i in tl.range(0, MAX_KV_PER_Q):
if i < n:
kv = tl.load(q2k_row + i * q2k_idx_k)
count_ptr = (
k2q_num_ptr + pid_b * k2q_num_b + pid_h * k2q_num_h + kv * k2q_num_k
)
pos = tl.atomic_add(count_ptr, 1)
tl.store(
k2q_idx_ptr
+ pid_b * k2q_idx_b
+ pid_h * k2q_idx_h
+ kv * k2q_idx_k
+ pos * k2q_idx_q,
pid_q,
)
def invert_indices(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
if q2k_idx.dim() != 4:
raise ValueError(
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
)
if q2k_num.dim() != 3:
raise ValueError(
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
)
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
raise RuntimeError("invert_indices requires CUDA tensors.")
B, H, Nq, Mk = q2k_idx.shape
if q2k_num.shape != (B, H, Nq):
raise ValueError(
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
f"[B, H, Nq] = {(B, H, Nq)}"
)
q2k_idx = q2k_idx.contiguous()
q2k_num = q2k_num.contiguous()
if q2k_idx.dtype != torch.int32:
q2k_idx = q2k_idx.to(torch.int32)
if q2k_num.dtype != torch.int32:
q2k_num = q2k_num.to(torch.int32)
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
# `Nq` is a tight upper bound on the compacted K->Q slots.
k2q_idx = torch.empty(
(B, H, num_kv_blocks, Nq),
dtype=torch.int32,
device=q2k_idx.device,
)
k2q_num = torch.zeros(
(B, H, num_kv_blocks),
dtype=torch.int32,
device=q2k_idx.device,
)
grid = (B, H, Nq)
_invert_indices_kernel[grid](
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
q2k_idx.stride(0),
q2k_idx.stride(1),
q2k_idx.stride(2),
q2k_idx.stride(3),
q2k_num.stride(0),
q2k_num.stride(1),
q2k_num.stride(2),
k2q_idx.stride(0),
k2q_idx.stride(1),
k2q_idx.stride(2),
k2q_idx.stride(3),
k2q_num.stride(0),
k2q_num.stride(1),
k2q_num.stride(2),
MAX_KV_PER_Q=Mk,
)
return k2q_idx, k2q_num