File size: 6,656 Bytes
7344bef | 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 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 | # References:
# https://github.com/facebookresearch/fairseq/blob/main/fairseq/modules/rotary_positional_embedding.py
import torch
import torch.nn as nn
from einops import rearrange
def broadcat(tensors, dim=-1):
num_tensors = len(tensors)
shape_lens = set(list(map(lambda t: len(t.shape), tensors)))
assert len(shape_lens) == 1, "tensors must all have the same number of dimensions"
shape_len = list(shape_lens)[0]
dim = (dim + shape_len) if dim < 0 else dim
dims = list(zip(*map(lambda t: list(t.shape), tensors)))
expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim]
assert all(
[*map(lambda t: len(set(t[1])) <= 2, expandable_dims)]
), "invalid dimensions for broadcastable concatentation"
max_dims = list(map(lambda t: (t[0], max(t[1])), expandable_dims))
expanded_dims = list(map(lambda t: (t[0], (t[1],) * num_tensors), max_dims))
expanded_dims.insert(dim, (dim, dims[dim]))
expandable_shapes = list(zip(*map(lambda t: t[1], expanded_dims)))
tensors = list(map(lambda t: t[0].expand(*t[1]), zip(tensors, expandable_shapes)))
return torch.cat(tensors, dim=dim)
def rotate_half(x):
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1)
return torch.stack((-x_imag, x_real), dim=-1).flatten(-2)
def apply_rotary_inplace(x, cos, sin):
out_shape = x.shape
x_pair = x.reshape(*out_shape[:-1], -1, 2)
if cos.shape[-1] == out_shape[-1]:
cos = cos[..., ::2]
sin = sin[..., ::2]
real = x_pair[..., 0]
imag = x_pair[..., 1]
scratch = real.clone()
real.mul_(cos).addcmul_(imag, sin, value=-1)
imag.mul_(cos).addcmul_(scratch, sin)
del scratch
return x_pair.reshape(out_shape)
class RotaryPositionalEmbedding(nn.Module):
def __init__(self,
head_dim,
cp_split_hw=None
):
"""Rotary positional embedding for 3D
Reference : https://blog.eleuther.ai/rotary-embeddings/
Paper: https://arxiv.org/pdf/2104.09864.pdf
Args:
dim: Dimension of embedding
base: Base value for exponential
"""
super().__init__()
self.head_dim = head_dim
assert self.head_dim % 8 == 0, 'Dim must be a multiply of 8 for 3D RoPE.'
self.cp_split_hw = cp_split_hw
# We take the assumption that the longest side of grid will not larger than 512, i.e, 512 * 8 = 4098 input pixels
self.base = 10000
self.freqs_dict = {}
def register_grid_size(self, grid_size, key_name, frame_index=None, num_ref_latents=None):
if key_name not in self.freqs_dict:
self.freqs_dict.update({
key_name: self.precompute_freqs_cis_3d(grid_size, frame_index, num_ref_latents)
})
def precompute_freqs_cis_3d(self, grid_size, frame_index=None, num_ref_latents=None):
num_frames, height, width = grid_size
dim_t = self.head_dim - 4 * (self.head_dim // 6)
dim_h = 2 * (self.head_dim // 6)
dim_w = 2 * (self.head_dim // 6)
cpu = torch.device("cpu")
freqs_t = 1.0 / (
self.base ** (torch.arange(0, dim_t, 2, device=cpu, dtype=torch.float32)[: (dim_t // 2)] / dim_t)
)
freqs_h = 1.0 / (
self.base ** (torch.arange(0, dim_h, 2, device=cpu, dtype=torch.float32)[: (dim_h // 2)] / dim_h)
)
freqs_w = 1.0 / (
self.base ** (torch.arange(0, dim_w, 2, device=cpu, dtype=torch.float32)[: (dim_w // 2)] / dim_w)
)
if frame_index is not None and num_ref_latents is not None:
grid_t = torch.concat(
[
torch.tensor([frame_index], device=cpu, dtype=torch.float32),
torch.arange(0, num_frames - num_ref_latents, device=cpu, dtype=torch.float32),
],
dim=0,
)
else:
grid_t = torch.arange(num_frames, device=cpu, dtype=torch.float32)
grid_h = torch.arange(height, device=cpu, dtype=torch.float32)
grid_w = torch.arange(width, device=cpu, dtype=torch.float32)
freqs_t = torch.einsum("..., f -> ... f", grid_t, freqs_t)
freqs_h = torch.einsum("..., f -> ... f", grid_h, freqs_h)
freqs_w = torch.einsum("..., f -> ... f", grid_w, freqs_w)
freqs = broadcat((freqs_t[:, None, None, :], freqs_h[None, :, None, :], freqs_w[None, None, :, :]), dim=-1)
# (T H W D)
freqs = rearrange(freqs, "T H W D -> (T H W) D")
return freqs
def forward(self, q, k, grid_size, frame_index=None, num_ref_latents=None):
"""3D RoPE.
Args:
query: [B, head, seq, head_dim]
key: [B, head, seq, head_dim]
Returns:
query and key with the same shape as input.
"""
key_name = '.'.join([str(i) for i in grid_size]) + f"-{str(frame_index)}-{str(num_ref_latents)}"
if key_name not in self.freqs_dict:
self.register_grid_size(grid_size, key_name, frame_index, num_ref_latents)
freqs = self.freqs_dict[key_name].to(device=q.device, dtype=torch.float32)
cos = freqs.cos().unsqueeze(0).unsqueeze(2)
sin = freqs.sin().unsqueeze(0).unsqueeze(2)
q = apply_rotary_inplace(q, cos, sin)
k = apply_rotary_inplace(k, cos, sin)
return q, k
class RotaryPositionalEmbedding1D(nn.Module):
def __init__(self,
head_dim
):
"""Rotary positional embedding for 1D
Args:
dim: Dimension of embedding
base: Base value for exponential
"""
super().__init__()
self.head_dim = head_dim
self.base = 10000
def precompute_freqs_cis_1d(self, pos_indices):
freqs = 1.0 / (self.base ** (torch.arange(0, self.head_dim, 2, device=pos_indices.device, dtype=torch.float32)[: (self.head_dim // 2)] / self.head_dim))
freqs = freqs.to(pos_indices.device)
freqs = torch.einsum("..., f -> ... f", pos_indices.float(), freqs)
return freqs
def forward(self, x, pos_indices):
"""1D RoPE.
Args:
query (torch.tensor): [B, seq, head, head_dim]
pos_indices (torch.tensor): [seq,]
Returns:
query with the same shape as input.
"""
freqs_cis = self.precompute_freqs_cis_1d(pos_indices)
freqs_cis = freqs_cis.float().to(x.device)
cos = freqs_cis.cos().unsqueeze(0).unsqueeze(2)
sin = freqs_cis.sin().unsqueeze(0).unsqueeze(2)
return apply_rotary_inplace(x, cos, sin)
|