| 266 | |
| 267 | class Rotary3DPositionEmbeddingMixin(BaseMixin): |
| 268 | def __init__( |
| 269 | self, |
| 270 | height, |
| 271 | width, |
| 272 | compressed_num_frames, |
| 273 | hidden_size, |
| 274 | hidden_size_head, |
| 275 | text_length, |
| 276 | theta=10000, |
| 277 | rot_v=False, |
| 278 | pnp=False, |
| 279 | learnable_pos_embed=False, |
| 280 | ): |
| 281 | super().__init__() |
| 282 | self.rot_v = rot_v |
| 283 | |
| 284 | dim_t = hidden_size_head // 4 |
| 285 | dim_h = hidden_size_head // 8 * 3 |
| 286 | dim_w = hidden_size_head // 8 * 3 |
| 287 | |
| 288 | # 'lang': |
| 289 | freqs_t = 1.0 / (theta ** (torch.arange(0, dim_t, 2)[: (dim_t // 2)].float() / dim_t)) |
| 290 | freqs_h = 1.0 / (theta ** (torch.arange(0, dim_h, 2)[: (dim_h // 2)].float() / dim_h)) |
| 291 | freqs_w = 1.0 / (theta ** (torch.arange(0, dim_w, 2)[: (dim_w // 2)].float() / dim_w)) |
| 292 | |
| 293 | grid_t = torch.arange(compressed_num_frames, dtype=torch.float32) |
| 294 | grid_h = torch.arange(height, dtype=torch.float32) |
| 295 | grid_w = torch.arange(width, dtype=torch.float32) |
| 296 | |
| 297 | freqs_t = torch.einsum("..., f -> ... f", grid_t, freqs_t) |
| 298 | freqs_h = torch.einsum("..., f -> ... f", grid_h, freqs_h) |
| 299 | freqs_w = torch.einsum("..., f -> ... f", grid_w, freqs_w) |
| 300 | |
| 301 | freqs_t = repeat(freqs_t, "... n -> ... (n r)", r=2) |
| 302 | freqs_h = repeat(freqs_h, "... n -> ... (n r)", r=2) |
| 303 | freqs_w = repeat(freqs_w, "... n -> ... (n r)", r=2) |
| 304 | |
| 305 | freqs = broadcat((freqs_t[:, None, None, :], freqs_h[None, :, None, :], freqs_w[None, None, :, :]), dim=-1) |
| 306 | # (T H W D) |
| 307 | |
| 308 | self.pnp = pnp |
| 309 | |
| 310 | if not self.pnp: |
| 311 | freqs = rearrange(freqs, "t h w d -> (t h w) d") |
| 312 | |
| 313 | freqs = freqs.contiguous() |
| 314 | freqs_sin = freqs.sin() |
| 315 | freqs_cos = freqs.cos() |
| 316 | self.register_buffer("freqs_sin", freqs_sin) |
| 317 | self.register_buffer("freqs_cos", freqs_cos) |
| 318 | |
| 319 | self.text_length = text_length |
| 320 | if learnable_pos_embed: |
| 321 | num_patches = height * width * compressed_num_frames + text_length |
| 322 | self.pos_embedding = nn.Parameter(torch.zeros(1, num_patches, int(hidden_size)), requires_grad=True) |
| 323 | else: |
| 324 | self.pos_embedding = None |
| 325 | |