MCPcopy Create free account
hub / github.com/MeiGen-AI/MultiTalk / forward

Method forward

wan/modules/attention.py:227–280  ·  view source on GitHub ↗
(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None)

Source from the content-addressed store, hash-verified

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
282class SingleStreamMutiAttention(SingleStreamAttention):
283 def __init__(

Callers 1

forwardMethod · 0.45

Calls 1

Tested by

no test coverage detected