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

Method forward

wan/modules/attention.py:316–393  ·  view source on GitHub ↗
(self, 
                x: torch.Tensor, 
                encoder_hidden_states: torch.Tensor, 
                shape=None, 
                x_ref_attn_map=None,
                human_num=None)

Source from the content-addressed store, hash-verified

314 self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim)
315
316 def forward(self,
317 x: torch.Tensor,
318 encoder_hidden_states: torch.Tensor,
319 shape=None,
320 x_ref_attn_map=None,
321 human_num=None) -> torch.Tensor:
322
323 encoder_hidden_states = encoder_hidden_states.squeeze(0)
324 if human_num == 1:
325 return super().forward(x, encoder_hidden_states, shape)
326
327 N_t, _, _ = shape
328 x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
329
330 # get q for hidden_state
331 B, N, C = x.shape
332 q = self.q_linear(x)
333 q_shape = (B, N, self.num_heads, self.head_dim)
334 q = q.view(q_shape).permute((0, 2, 1, 3))
335
336 if self.qk_norm:
337 q = self.q_norm(q)
338
339
340 max_values = x_ref_attn_map.max(1).values[:, None, None]
341 min_values = x_ref_attn_map.min(1).values[:, None, None]
342 max_min_values = torch.cat([max_values, min_values], dim=2)
343
344 human1_max_value, human1_min_value = max_min_values[0, :, 0].max(), max_min_values[0, :, 1].min()
345 human2_max_value, human2_min_value = max_min_values[1, :, 0].max(), max_min_values[1, :, 1].min()
346
347 human1 = normalize_and_scale(x_ref_attn_map[0], (human1_min_value, human1_max_value), (self.rope_h1[0], self.rope_h1[1]))
348 human2 = normalize_and_scale(x_ref_attn_map[1], (human2_min_value, human2_max_value), (self.rope_h2[0], self.rope_h2[1]))
349 back = torch.full((x_ref_attn_map.size(1),), self.rope_bak, dtype=human1.dtype).to(human1.device)
350 max_indices = x_ref_attn_map.argmax(dim=0)
351 normalized_map = torch.stack([human1, human2, back], dim=1)
352 normalized_pos = normalized_map[range(x_ref_attn_map.size(1)), max_indices] # N
353
354 q = rearrange(q, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
355 q = self.rope_1d(q, normalized_pos)
356 q = rearrange(q, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
357
358 _, N_a, _ = encoder_hidden_states.shape
359 encoder_kv = self.kv_linear(encoder_hidden_states)
360 encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
361 encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
362 encoder_k, encoder_v = encoder_kv.unbind(0)
363
364 if self.qk_norm:
365 encoder_k = self.add_k_norm(encoder_k)
366
367
368 per_frame = torch.zeros(N_a, dtype=encoder_k.dtype).to(encoder_k.device)
369 per_frame[:per_frame.size(0)//2] = (self.rope_h1[0] + self.rope_h1[1]) / 2
370 per_frame[per_frame.size(0)//2:] = (self.rope_h2[0] + self.rope_h2[1]) / 2
371 encoder_pos = torch.concat([per_frame]*N_t, dim=0)
372 encoder_k = rearrange(encoder_k, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
373 encoder_k = self.rope_1d(encoder_k, encoder_pos)

Callers

nothing calls this directly

Calls 2

normalize_and_scaleFunction · 0.85
forwardMethod · 0.45

Tested by

no test coverage detected