Add initial support for Minimax H3 VSA sparse attention
This commit is contained in:
@@ -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,
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -1 +1 @@
|
|||||||
from .minimax_h3 import MinimaxH3Model, MinimaxH3Ref2VAModel
|
from .minimax_h3 import MinimaxH3Model, MinimaxH3Ref2VAModel, MinimaxH3FastModel
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
343
extensions_built_in/diffusion_models/minimax_h3/src/vsa.py
Normal file
343
extensions_built_in/diffusion_models/minimax_h3/src/vsa.py
Normal 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)
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user