Ref2VA/video_vae/flash.py
5.6 KB · 179 lines · python Raw
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