| 9 | |
| 10 | |
| 11 | class CogPatchify(torch.nn.Module): |
| 12 | def __init__(self, dim_in, dim_out, patch_size) -> None: |
| 13 | super().__init__() |
| 14 | self.proj = torch.nn.Conv3d(dim_in, dim_out, kernel_size=(1, patch_size, patch_size), stride=(1, patch_size, patch_size)) |
| 15 | |
| 16 | def forward(self, hidden_states): |
| 17 | hidden_states = self.proj(hidden_states) |
| 18 | hidden_states = rearrange(hidden_states, "B C T H W -> B (T H W) C") |
| 19 | return hidden_states |
| 20 | |
| 21 | |
| 22 |