MCPcopy Create free account
hub / github.com/WeChatCV/WeVisionOne / VisionRotaryEmbeddingFast

Class VisionRotaryEmbeddingFast

WeVisionOne/backbone/eva02/det/utils.py:309–346  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

307
308
309class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected