from math import log, pi import torch from torch import nn, einsum import torch.nn.functional as F from einops import rearrange, repeat def rotate_every_two(x): x = rearrange(x, '... (d j) -> ... d j', j = 2) x1, x2 = x.unbind(dim = -1) x = torch.stack((-x2, x1), dim = -1) return rearrange(x, '... d j -> ... (d j)') def apply_rot_emb(q, k, rot_emb): sin, cos = rot_emb rot_dim = sin.shape[-1] (q, q_pass), (k, k_pass) = map(lambda t: (t[..., :rot_dim], t[..., rot_dim:]), (q, k)) q, k = map(lambda t: t * cos + rotate_every_two(t) * sin, (q, k)) q, k = map(lambda t: torch.cat(t, dim = -1), ((q, q_pass), (k, k_pass))) return q, k class AxialRotaryEmbedding(nn.Module): def __init__(self, dim, max_freq = 10): super().__init__() self.dim = dim scales = torch.logspace(0., log(max_freq / 2) / log(2), self.dim // 4, base = 2) self.register_buffer('scales', scales) def forward(self, h, w, device): scales = rearrange(self.scales, '... -> () ...') scales = scales.to(device) h_seq = torch.linspace(-1., 1., steps = h, device = device) h_seq = h_seq.unsqueeze(-1) w_seq = torch.linspace(-1., 1., steps = w, device = device) w_seq = w_seq.unsqueeze(-1) h_seq = h_seq * scales * pi w_seq = w_seq * scales * pi x_sinu = repeat(h_seq, 'i d -> i j d', j = w) y_sinu = repeat(w_seq, 'j d -> i j d', i = h) sin = torch.cat((x_sinu.sin(), y_sinu.sin()), dim = -1) cos = torch.cat((x_sinu.cos(), y_sinu.cos()), dim = -1) sin, cos = map(lambda t: rearrange(t, 'i j d -> (i j) d'), (sin, cos)) sin, cos = map(lambda t: repeat(t, 'n d -> () n (d j)', j = 2), (sin, cos)) return sin, cos class RotaryEmbedding(nn.Module): def __init__(self, dim): super().__init__() inv_freqs = 1. / (10000 ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer('inv_freqs', inv_freqs) def forward(self, n, device): seq = torch.arange(n, device = device) freqs = einsum('i, j -> i j', seq, self.inv_freqs) freqs = torch.cat((freqs, freqs), dim = -1) freqs = rearrange(freqs, 'n d -> () n d') return freqs.sin(), freqs.cos()