(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None)
| 225 | self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() |
| 226 | |
| 227 | def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor: |
| 228 | |
| 229 | N_t, N_h, N_w = shape |
| 230 | if not enable_sp: |
| 231 | x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) |
| 232 | |
| 233 | # get q for hidden_state |
| 234 | B, N, C = x.shape |
| 235 | q = self.q_linear(x) |
| 236 | q_shape = (B, N, self.num_heads, self.head_dim) |
| 237 | q = q.view(q_shape).permute((0, 2, 1, 3)) |
| 238 | |
| 239 | if self.qk_norm: |
| 240 | q = self.q_norm(q) |
| 241 | |
| 242 | # get kv from encoder_hidden_states |
| 243 | _, N_a, _ = encoder_hidden_states.shape |
| 244 | encoder_kv = self.kv_linear(encoder_hidden_states) |
| 245 | encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim) |
| 246 | encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4)) |
| 247 | encoder_k, encoder_v = encoder_kv.unbind(0) |
| 248 | |
| 249 | if self.qk_norm: |
| 250 | encoder_k = self.add_k_norm(encoder_k) |
| 251 | |
| 252 | |
| 253 | q = rearrange(q, "B H M K -> B M H K") |
| 254 | encoder_k = rearrange(encoder_k, "B H M K -> B M H K") |
| 255 | encoder_v = rearrange(encoder_v, "B H M K -> B M H K") |
| 256 | |
| 257 | if enable_sp: |
| 258 | # context parallel |
| 259 | sp_size = get_sequence_parallel_world_size() |
| 260 | sp_rank = get_sequence_parallel_rank() |
| 261 | visual_seqlen, _ = split_token_counts_and_frame_ids(N_t, N_h * N_w, sp_size, sp_rank) |
| 262 | assert kv_seq is not None, f"kv_seq should not be None." |
| 263 | attn_bias = xformers.ops.fmha.attn_bias.BlockDiagonalMask.from_seqlens(visual_seqlen, kv_seq) |
| 264 | else: |
| 265 | attn_bias = None |
| 266 | x = xformers.ops.memory_efficient_attention(q, encoder_k, encoder_v, attn_bias=attn_bias, op=None,) |
| 267 | x = rearrange(x, "B M H K -> B H M K") |
| 268 | |
| 269 | # linear transform |
| 270 | x_output_shape = (B, N, C) |
| 271 | x = x.transpose(1, 2) |
| 272 | x = x.reshape(x_output_shape) |
| 273 | x = self.proj(x) |
| 274 | x = self.proj_drop(x) |
| 275 | |
| 276 | if not enable_sp: |
| 277 | # reshape x to origin shape |
| 278 | x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) |
| 279 | |
| 280 | return x |
| 281 | |
| 282 | class SingleStreamMutiAttention(SingleStreamAttention): |
| 283 | def __init__( |
no test coverage detected