| |
| |
| 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): |
| |
| |
| 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) |
|
|
|
|
| |
| |
| |
|
|
|
|
| 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 |
|
|
|
|
|
|
|
|
|
|
|
|
| |
| |
| |
|
|
|
|
| @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) |
|
|