(self,
x: torch.Tensor,
encoder_hidden_states: torch.Tensor,
shape=None,
x_ref_attn_map=None,
human_num=None)
| 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) |
nothing calls this directly
no test coverage detected