FL2VA/video_vae/flash.py
| 1 | # SPDX-License-Identifier: Apache-2.0 |
| 2 | # Torch-native attention implemented with PyTorch SDPA instead of FA4/CUTLASS. |
| 3 | import os |
| 4 | from contextlib import nullcontext |
| 5 | |
| 6 | import torch |
| 7 | import torch.nn.functional as F |
| 8 | |
| 9 | |
| 10 | _BLOCK_CAUSAL_MASK_MOD_CACHE = {} |
| 11 | |
| 12 | |
| 13 | |
| 14 | |
| 15 | def _as_bool_mask(mask, *, device): |
| 16 | if not isinstance(mask, torch.Tensor): |
| 17 | mask = torch.as_tensor(mask, device=device) |
| 18 | return mask.to(device=device, dtype=torch.bool) |
| 19 | |
| 20 | |
| 21 | def _ensure_nonempty_rows(mask): |
| 22 | if mask.numel() == 0 or mask.shape[-1] == 0: |
| 23 | return mask |
| 24 | empty = ~mask.any(dim=-1) |
| 25 | if empty.any(): |
| 26 | mask = mask.clone() |
| 27 | mask[..., 0] |= empty |
| 28 | return mask |
| 29 | |
| 30 | |
| 31 | def _sdpa_kernel_context(): |
| 32 | backend_name = os.environ.get("MINIMAX_H3_TORCH_SDPA_BACKEND", "auto").lower() |
| 33 | if backend_name in {"", "auto", "default"}: |
| 34 | return nullcontext() |
| 35 | |
| 36 | from torch.nn.attention import SDPBackend, sdpa_kernel |
| 37 | |
| 38 | backends = { |
| 39 | "math": SDPBackend.MATH, |
| 40 | "flash": SDPBackend.FLASH_ATTENTION, |
| 41 | "flash_attention": SDPBackend.FLASH_ATTENTION, |
| 42 | "efficient": SDPBackend.EFFICIENT_ATTENTION, |
| 43 | "mem_efficient": SDPBackend.EFFICIENT_ATTENTION, |
| 44 | "cudnn": SDPBackend.CUDNN_ATTENTION, |
| 45 | "cudnn_attention": SDPBackend.CUDNN_ATTENTION, |
| 46 | } |
| 47 | if backend_name not in backends: |
| 48 | raise ValueError( |
| 49 | "MINIMAX_H3_TORCH_SDPA_BACKEND must be one of " |
| 50 | f"{sorted([*backends, 'auto', 'default'])}, got {backend_name!r}" |
| 51 | ) |
| 52 | return sdpa_kernel(backends=[backends[backend_name]]) |
| 53 | |
| 54 | |
| 55 | def _sdpa_attention(query, key, value, causal=False, attn_mask=None): |
| 56 | # query/key/value arrive as [B, S, H, D]; PyTorch SDPA expects |
| 57 | # [B, H, S, D]. |
| 58 | q = query.transpose(1, 2) |
| 59 | k = key.transpose(1, 2) |
| 60 | v = value.transpose(1, 2) |
| 61 | if attn_mask is not None and attn_mask.dim() == 3: |
| 62 | attn_mask = attn_mask.unsqueeze(0) |
| 63 | with _sdpa_kernel_context(): |
| 64 | out = F.scaled_dot_product_attention( |
| 65 | q, |
| 66 | k, |
| 67 | v, |
| 68 | attn_mask=attn_mask, |
| 69 | dropout_p=0.0, |
| 70 | is_causal=causal, |
| 71 | ) |
| 72 | return out.transpose(1, 2).nan_to_num(0.0) |
| 73 | |
| 74 | |
| 75 | def _mask_mod_to_dense(mask_mod, batch, heads, q_len, kv_len, device, aux_tensors=None): |
| 76 | q_idx = torch.arange(q_len, device=device).view(q_len, 1) |
| 77 | kv_idx = torch.arange(kv_len, device=device).view(1, kv_len) |
| 78 | dense = torch.empty((batch, heads, q_len, kv_len), dtype=torch.bool, device=device) |
| 79 | for b in range(batch): |
| 80 | b_idx = torch.tensor(b, device=device) |
| 81 | for h in range(heads): |
| 82 | h_idx = torch.tensor(h, device=device) |
| 83 | mask = mask_mod(b_idx, h_idx, q_idx, kv_idx, None, aux_tensors) |
| 84 | dense[b, h] = _as_bool_mask(mask, device=device) |
| 85 | return _ensure_nonempty_rows(dense) |
| 86 | |
| 87 | |
| 88 | ######################################################### |
| 89 | # Block causal attention |
| 90 | ######################################################### |
| 91 | |
| 92 | |
| 93 | def make_block_causal_mask_mod(num_tokens, block_size, num_special=0, suffix=False): |
| 94 | if num_tokens < 0: |
| 95 | raise ValueError(f"num_tokens must be non-negative, got {num_tokens}") |
| 96 | if block_size <= 0: |
| 97 | raise ValueError(f"block_size must be positive, got {block_size}") |
| 98 | if num_special < 0: |
| 99 | raise ValueError(f"num_special must be non-negative, got {num_special}") |
| 100 | |
| 101 | cache_key = (num_tokens, block_size, num_special, suffix) |
| 102 | if cache_key in _BLOCK_CAUSAL_MASK_MOD_CACHE: |
| 103 | return _BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key] |
| 104 | |
| 105 | if suffix: |
| 106 | |
| 107 | def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors): |
| 108 | del b, h, seqlen_info, aux_tensors |
| 109 | q_is_special = q_idx >= num_tokens |
| 110 | kv_is_special = kv_idx >= num_tokens |
| 111 | return q_is_special | kv_is_special | ( |
| 112 | q_idx // block_size >= kv_idx // block_size |
| 113 | ) |
| 114 | |
| 115 | else: |
| 116 | |
| 117 | def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors): |
| 118 | del b, h, seqlen_info, aux_tensors |
| 119 | q_is_special = q_idx < num_special |
| 120 | kv_is_special = kv_idx < num_special |
| 121 | q_block_idx = (q_idx - num_special) // block_size |
| 122 | kv_block_idx = (kv_idx - num_special) // block_size |
| 123 | return q_is_special | kv_is_special | (q_block_idx >= kv_block_idx) |
| 124 | |
| 125 | mask_mod.block_sparse_cache_key = ( |
| 126 | "block_causal", |
| 127 | num_tokens, |
| 128 | block_size, |
| 129 | num_special, |
| 130 | suffix, |
| 131 | ) |
| 132 | _BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key] = mask_mod |
| 133 | return mask_mod |
| 134 | |
| 135 | |
| 136 | |
| 137 | |
| 138 | |
| 139 | |
| 140 | ######################################################### |
| 141 | # Public entry point |
| 142 | ######################################################### |
| 143 | |
| 144 | |
| 145 | @torch.compiler.disable |
| 146 | def flash_attn( |
| 147 | query: torch.Tensor, |
| 148 | key: torch.Tensor, |
| 149 | value: torch.Tensor, |
| 150 | causal: bool = False, |
| 151 | mask_mod=None, |
| 152 | block_sparse=None, |
| 153 | aux_tensors=None, |
| 154 | ) -> torch.Tensor: |
| 155 | use_masked = mask_mod is not None or block_sparse is not None |
| 156 | |
| 157 | if block_sparse is not None and mask_mod is None: |
| 158 | raise ValueError("block_sparse requires mask_mod") |
| 159 | if causal and mask_mod is not None: |
| 160 | raise ValueError("causal must be encoded in mask_mod when using masked attention") |
| 161 | if aux_tensors is not None and not use_masked: |
| 162 | raise ValueError("aux_tensors is only supported with masked attention") |
| 163 | |
| 164 | if use_masked: |
| 165 | batch, q_len, heads, _ = query.shape |
| 166 | kv_len = key.shape[1] |
| 167 | dense_mask = _mask_mod_to_dense( |
| 168 | mask_mod, |
| 169 | batch, |
| 170 | heads, |
| 171 | q_len, |
| 172 | kv_len, |
| 173 | query.device, |
| 174 | aux_tensors=aux_tensors, |
| 175 | ) |
| 176 | return _sdpa_attention(query, key, value, attn_mask=dense_mask) |
| 177 | |
| 178 | return _sdpa_attention(query, key, value, causal=causal) |
| 179 | |