(self, query_layer, key_layer, value_layer, attention_mask)
| 189 | self.attention_dropout = config.attention_dropout |
| 190 | |
| 191 | def forward(self, query_layer, key_layer, value_layer, attention_mask): |
| 192 | seqlen_q, batch_size = query_layer.shape[0], query_layer.shape[1] |
| 193 | seqlen_k = key_layer.shape[0] |
| 194 | query_layer, key_layer, value_layer = [rearrange(x, 's b ... -> (b s) ...') for x in [query_layer, key_layer, value_layer]] |
| 195 | # DO flash_attn_varlen_func |
| 196 | if attention_mask is None or attention_mask.ndim != 1: |
| 197 | cu_seqlens_q = torch.arange(0, (batch_size + 1) * seqlen_q, step=seqlen_q, dtype=torch.int32, |
| 198 | device=query_layer.device) |
| 199 | else: |
| 200 | assert seqlen_q == seqlen_k |
| 201 | cu_seqlens_q = attention_mask |
| 202 | if self.training: |
| 203 | assert seqlen_k == seqlen_q |
| 204 | is_causal = True |
| 205 | cu_seqlens_k = cu_seqlens_q |
| 206 | else: |
| 207 | is_causal = seqlen_q == seqlen_k |
| 208 | cu_seqlens_k = torch.arange(0, (batch_size + 1) * seqlen_k, step=seqlen_k, dtype=torch.int32, |
| 209 | device=query_layer.device) if not is_causal else cu_seqlens_q |
| 210 | self.attention_dropout = 0 |
| 211 | context_layer = flash_attn_unpadded_func( |
| 212 | query_layer, key_layer, value_layer, cu_seqlens_q, cu_seqlens_k, seqlen_q, seqlen_k, |
| 213 | self.attention_dropout, |
| 214 | softmax_scale=1.0 / self.norm_factor, causal=is_causal |
| 215 | ) |
| 216 | context_layer = rearrange(context_layer, '(b s) ... -> s b ...', b=batch_size) |
| 217 | new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,) |
| 218 | context_layer = context_layer.reshape(*new_context_layer_shape) |
| 219 | return context_layer |
| 220 | |
| 221 | |
| 222 | class SelfAttention(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected