Patchifies and embeds the input tensor(s). Args: x (List[torch.Tensor] | torch.Tensor): The input tensor(s) to be patchified and embedded. Returns: Tuple[torch.Tensor, torch.Tensor, List[Tuple[int, int]], torch.Tensor]: A tuple containing the patchi
(self, x, freqs_cis)
| 590 | ) |
| 591 | |
| 592 | def forward(self, x, freqs_cis): |
| 593 | """ |
| 594 | Patchifies and embeds the input tensor(s). |
| 595 | |
| 596 | Args: |
| 597 | x (List[torch.Tensor] | torch.Tensor): The input tensor(s) to be patchified and embedded. |
| 598 | |
| 599 | Returns: |
| 600 | Tuple[torch.Tensor, torch.Tensor, List[Tuple[int, int]], torch.Tensor]: A tuple containing the patchified |
| 601 | and embedded tensor(s), the mask indicating the valid patches, the original image size(s), and the |
| 602 | frequency tensor(s). |
| 603 | """ |
| 604 | freqs_cis = freqs_cis.to(x[0].device) |
| 605 | patch_height = patch_width = self.patch_size |
| 606 | batch_size, channel, height, width = x.size() |
| 607 | height_tokens, width_tokens = height // patch_height, width // patch_width |
| 608 | |
| 609 | x = x.view(batch_size, channel, height_tokens, patch_height, width_tokens, patch_width).permute( |
| 610 | 0, 2, 4, 1, 3, 5 |
| 611 | ) |
| 612 | x = x.flatten(3) |
| 613 | x = self.proj(x) |
| 614 | x = x.flatten(1, 2) |
| 615 | |
| 616 | mask = torch.ones(x.shape[0], x.shape[1], dtype=torch.int32, device=x.device) |
| 617 | |
| 618 | return ( |
| 619 | x, |
| 620 | mask, |
| 621 | [(height, width)] * batch_size, |
| 622 | freqs_cis[:height_tokens, :width_tokens].flatten(0, 1).unsqueeze(0), |
| 623 | ) |
| 624 | |
| 625 | |
| 626 | class CogVideoXPatchEmbed(nn.Module): |