Green-eyedDevil's picture
Duplicate from MiniMaxAI/MiniMax-H3
fae01f4
Raw
History Blame Contribute Delete
5.76 kB
# SPDX-License-Identifier: Apache-2.0
# Torch-native attention implemented with PyTorch SDPA instead of FA4/CUTLASS.
import os
from contextlib import nullcontext
import torch
import torch.nn.functional as F
_BLOCK_CAUSAL_MASK_MOD_CACHE = {}
def _as_bool_mask(mask, *, device):
if not isinstance(mask, torch.Tensor):
mask = torch.as_tensor(mask, device=device)
return mask.to(device=device, dtype=torch.bool)
def _ensure_nonempty_rows(mask):
if mask.numel() == 0 or mask.shape[-1] == 0:
return mask
empty = ~mask.any(dim=-1)
if empty.any():
mask = mask.clone()
mask[..., 0] |= empty
return mask
def _sdpa_kernel_context():
backend_name = os.environ.get("MINIMAX_H3_TORCH_SDPA_BACKEND", "auto").lower()
if backend_name in {"", "auto", "default"}:
return nullcontext()
from torch.nn.attention import SDPBackend, sdpa_kernel
backends = {
"math": SDPBackend.MATH,
"flash": SDPBackend.FLASH_ATTENTION,
"flash_attention": SDPBackend.FLASH_ATTENTION,
"efficient": SDPBackend.EFFICIENT_ATTENTION,
"mem_efficient": SDPBackend.EFFICIENT_ATTENTION,
"cudnn": SDPBackend.CUDNN_ATTENTION,
"cudnn_attention": SDPBackend.CUDNN_ATTENTION,
}
if backend_name not in backends:
raise ValueError(
"MINIMAX_H3_TORCH_SDPA_BACKEND must be one of "
f"{sorted([*backends, 'auto', 'default'])}, got {backend_name!r}"
)
return sdpa_kernel(backends=[backends[backend_name]])
def _sdpa_attention(query, key, value, causal=False, attn_mask=None):
# query/key/value arrive as [B, S, H, D]; PyTorch SDPA expects
# [B, H, S, D].
q = query.transpose(1, 2)
k = key.transpose(1, 2)
v = value.transpose(1, 2)
if attn_mask is not None and attn_mask.dim() == 3:
attn_mask = attn_mask.unsqueeze(0)
with _sdpa_kernel_context():
out = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=attn_mask,
dropout_p=0.0,
is_causal=causal,
)
return out.transpose(1, 2).nan_to_num(0.0)
def _mask_mod_to_dense(mask_mod, batch, heads, q_len, kv_len, device, aux_tensors=None):
q_idx = torch.arange(q_len, device=device).view(q_len, 1)
kv_idx = torch.arange(kv_len, device=device).view(1, kv_len)
dense = torch.empty((batch, heads, q_len, kv_len), dtype=torch.bool, device=device)
for b in range(batch):
b_idx = torch.tensor(b, device=device)
for h in range(heads):
h_idx = torch.tensor(h, device=device)
mask = mask_mod(b_idx, h_idx, q_idx, kv_idx, None, aux_tensors)
dense[b, h] = _as_bool_mask(mask, device=device)
return _ensure_nonempty_rows(dense)
#########################################################
# Block causal attention
#########################################################
def make_block_causal_mask_mod(num_tokens, block_size, num_special=0, suffix=False):
if num_tokens < 0:
raise ValueError(f"num_tokens must be non-negative, got {num_tokens}")
if block_size <= 0:
raise ValueError(f"block_size must be positive, got {block_size}")
if num_special < 0:
raise ValueError(f"num_special must be non-negative, got {num_special}")
cache_key = (num_tokens, block_size, num_special, suffix)
if cache_key in _BLOCK_CAUSAL_MASK_MOD_CACHE:
return _BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key]
if suffix:
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
del b, h, seqlen_info, aux_tensors
q_is_special = q_idx >= num_tokens
kv_is_special = kv_idx >= num_tokens
return q_is_special | kv_is_special | (
q_idx // block_size >= kv_idx // block_size
)
else:
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
del b, h, seqlen_info, aux_tensors
q_is_special = q_idx < num_special
kv_is_special = kv_idx < num_special
q_block_idx = (q_idx - num_special) // block_size
kv_block_idx = (kv_idx - num_special) // block_size
return q_is_special | kv_is_special | (q_block_idx >= kv_block_idx)
mask_mod.block_sparse_cache_key = (
"block_causal",
num_tokens,
block_size,
num_special,
suffix,
)
_BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key] = mask_mod
return mask_mod
#########################################################
# Public entry point
#########################################################
@torch.compiler.disable
def flash_attn(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
causal: bool = False,
mask_mod=None,
block_sparse=None,
aux_tensors=None,
) -> torch.Tensor:
use_masked = mask_mod is not None or block_sparse is not None
if block_sparse is not None and mask_mod is None:
raise ValueError("block_sparse requires mask_mod")
if causal and mask_mod is not None:
raise ValueError("causal must be encoded in mask_mod when using masked attention")
if aux_tensors is not None and not use_masked:
raise ValueError("aux_tensors is only supported with masked attention")
if use_masked:
batch, q_len, heads, _ = query.shape
kv_len = key.shape[1]
dense_mask = _mask_mod_to_dense(
mask_mod,
batch,
heads,
q_len,
kv_len,
query.device,
aux_tensors=aux_tensors,
)
return _sdpa_attention(query, key, value, attn_mask=dense_mask)
return _sdpa_attention(query, key, value, causal=causal)