File size: 2,286 Bytes
9e14838 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 | 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() |