| 307 | |
| 308 | |
| 309 | class VisionRotaryEmbeddingFast(nn.Module): |
| 310 | def __init__( |
| 311 | self, |
| 312 | dim, |
| 313 | pt_seq_len=16, |
| 314 | ft_seq_len=None, |
| 315 | custom_freqs = None, |
| 316 | freqs_for = 'lang', |
| 317 | theta = 10000, |
| 318 | max_freq = 10, |
| 319 | num_freqs = 1, |
| 320 | ): |
| 321 | super().__init__() |
| 322 | if custom_freqs: |
| 323 | freqs = custom_freqs |
| 324 | elif freqs_for == 'lang': |
| 325 | freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) |
| 326 | elif freqs_for == 'pixel': |
| 327 | freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi |
| 328 | elif freqs_for == 'constant': |
| 329 | freqs = torch.ones(num_freqs).float() |
| 330 | else: |
| 331 | raise ValueError(f'unknown modality {freqs_for}') |
| 332 | |
| 333 | if ft_seq_len is None: ft_seq_len = pt_seq_len |
| 334 | t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len |
| 335 | |
| 336 | freqs = torch.einsum('..., f -> ... f', t, freqs) |
| 337 | freqs = repeat(freqs, '... n -> ... (n r)', r = 2) |
| 338 | freqs = broadcat((freqs[:, None, :], freqs[None, :, :]), dim = -1) |
| 339 | |
| 340 | freqs_cos = freqs.cos().view(-1, freqs.shape[-1]) |
| 341 | freqs_sin = freqs.sin().view(-1, freqs.shape[-1]) |
| 342 | |
| 343 | self.register_buffer("freqs_cos", freqs_cos) |
| 344 | self.register_buffer("freqs_sin", freqs_sin) |
| 345 | |
| 346 | def forward(self, t): return t * self.freqs_cos + rotate_half(t) * self.freqs_sin |
| 347 | |