(self, x, attn_mask: Optional[torch.Tensor] = None)
| 219 | self.out_drop = nn.Dropout(proj_drop) |
| 220 | |
| 221 | def forward(self, x, attn_mask: Optional[torch.Tensor] = None): |
| 222 | if self.batch_first: |
| 223 | x = x.transpose(0, 1) |
| 224 | |
| 225 | L, N, C = x.shape |
| 226 | q, k, v = F.linear(x, self.in_proj_weight, self.in_proj_bias).chunk(3, dim=-1) |
| 227 | q = q.reshape(L, N * self.num_heads, -1).transpose(0, 1) |
| 228 | k = k.reshape(L, N * self.num_heads, -1).transpose(0, 1) |
| 229 | v = v.reshape(L, N * self.num_heads, -1).transpose(0, 1) |
| 230 | |
| 231 | if attn_mask is not None and attn_mask.dtype == torch.bool: |
| 232 | new_attn_mask = torch.zeros_like(attn_mask, dtype=q.dtype) |
| 233 | new_attn_mask.masked_fill_(attn_mask, float("-inf")) |
| 234 | attn_mask = new_attn_mask |
| 235 | |
| 236 | if self.logit_scale is not None: |
| 237 | attn = torch.bmm(F.normalize(q, dim=-1), F.normalize(k, dim=-1).transpose(-1, -2)) |
| 238 | logit_scale = torch.clamp(self.logit_scale, max=self.logit_scale_max).exp() |
| 239 | attn = attn.view(N, self.num_heads, L, L) * logit_scale |
| 240 | attn = attn.view(-1, L, L) |
| 241 | if attn_mask is not None: |
| 242 | attn = attn + attn_mask |
| 243 | attn = attn.softmax(dim=-1) |
| 244 | attn = self.attn_drop(attn) |
| 245 | x = torch.bmm(attn, v) |
| 246 | else: |
| 247 | if self.use_fsdpa: |
| 248 | x = F.scaled_dot_product_attention( |
| 249 | q, k, v, |
| 250 | attn_mask=attn_mask, |
| 251 | dropout_p=self.attn_drop.p if self.training else 0., |
| 252 | ) |
| 253 | else: |
| 254 | q = q * self.scale |
| 255 | attn = torch.bmm(q, k.transpose(-1, -2)) |
| 256 | if attn_mask is not None: |
| 257 | attn += attn_mask |
| 258 | attn = attn.softmax(dim=-1) |
| 259 | attn = self.attn_drop(attn) |
| 260 | x = torch.bmm(attn, v) |
| 261 | |
| 262 | if self.head_scale is not None: |
| 263 | x = x.view(N, self.num_heads, L, C) * self.head_scale |
| 264 | x = x.view(-1, L, C) |
| 265 | |
| 266 | x = x.transpose(0, 1).reshape(L, N, C) |
| 267 | |
| 268 | if self.batch_first: |
| 269 | x = x.transpose(0, 1) |
| 270 | |
| 271 | x = self.out_proj(x) |
| 272 | x = self.out_drop(x) |
| 273 | return x |
| 274 | |
| 275 | |
| 276 | class AttentionalPooler(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected