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)