yezdata commited on
Commit
599b1e0
·
verified ·
1 Parent(s): 32a5b3b

fix rope freq saving into model state dict

Browse files
Files changed (1) hide show
  1. rope_embeddings.py +88 -57
rope_embeddings.py CHANGED
@@ -10,39 +10,38 @@ from typing import Literal
10
  def exists(val):
11
  return val is not None
12
 
 
13
  def default(val, d):
14
  return val if exists(val) else d
15
 
 
16
  def broadcat(tensors, dim=-1):
17
  broadcasted_tensors = broadcast_tensors(*tensors)
18
  return torch.cat(broadcasted_tensors, dim=dim)
19
 
 
20
  def slice_at_dim(t, dim_slice: slice, *, dim):
21
- dim += (t.ndim if dim < 0 else 0)
22
  colons = [slice(None)] * t.ndim
23
  colons[dim] = dim_slice
24
  return t[tuple(colons)]
25
 
 
26
  def rotate_half(x):
27
  orig_shape = x.shape
28
  d_head = orig_shape[-1]
29
  x = x.view(*orig_shape[:-1], d_head // 2, 2)
30
-
31
  x1 = x[..., 0]
32
  x2 = x[..., 1]
33
-
34
  res = torch.stack((-x2, x1), dim=-1)
35
  return res.view(*orig_shape)
36
 
37
 
38
- @autocast('cuda', enabled=False)
39
  def apply_rotary_emb(
40
- freqs,
41
- t,
42
- start_index=0,
43
- scale=1.,
44
- seq_dim=-2,
45
- freqs_seq_dim=None
46
  ):
47
  dtype = t.dtype
48
 
@@ -57,21 +56,25 @@ def apply_rotary_emb(
57
  rot_dim = freqs.shape[-1]
58
  end_index = start_index + rot_dim
59
 
60
- assert rot_dim <= t.shape[-1], f'feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}'
 
 
61
 
62
  t_left = t[..., :start_index]
63
  t_middle = t[..., start_index:end_index]
64
  t_right = t[..., end_index:]
65
 
66
- t_transformed = (t_middle * freqs.cos() * scale) + (rotate_half(t_middle) * freqs.sin() * scale)
67
-
 
 
68
  out = torch.cat((t_left, t_transformed, t_right), dim=-1)
69
  return out.type(dtype)
70
 
71
 
72
  def apply_learned_rotations(rotations, t, start_index=0, freq_ranges=None):
73
  if exists(freq_ranges):
74
- rotations = torch.einsum('..., f -> ... f', rotations, freq_ranges)
75
  rotations = rotations.reshape(*rotations.shape[:-2], -1)
76
 
77
  rotations = rotations.repeat_interleave(2, dim=-1)
@@ -83,18 +86,18 @@ class RotaryEmbedding(Module):
83
  self,
84
  dim,
85
  custom_freqs: Tensor | None = None,
86
- freqs_for: Literal['lang', 'pixel', 'constant'] = 'lang',
87
- theta = 10000,
88
- max_freq = 10,
89
- num_freqs = 1,
90
- learned_freq = False,
91
- use_xpos = False,
92
- xpos_scale_base = 512,
93
- interpolate_factor = 1.,
94
- theta_rescale_factor = 1.,
95
- seq_before_head_dim = False,
96
- cache_if_possible = True,
97
- cache_max_seq_len = 8192
98
  ):
99
  super().__init__()
100
 
@@ -103,28 +106,35 @@ class RotaryEmbedding(Module):
103
 
104
  if exists(custom_freqs):
105
  freqs = custom_freqs
106
- elif freqs_for == 'lang':
107
- freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim))
108
- elif freqs_for == 'pixel':
109
- freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi
110
- elif freqs_for == 'constant':
 
 
111
  freqs = torch.ones(num_freqs).float()
112
 
113
  self.cache_if_possible = cache_if_possible
114
  self.cache_max_seq_len = cache_max_seq_len
115
 
116
- self.register_buffer('cached_freqs', torch.zeros(cache_max_seq_len, dim), persistent=False)
 
 
117
  self.cached_freqs_seq_len = 0
118
 
119
- self.freqs = nn.Parameter(freqs, requires_grad=learned_freq)
120
  self.learned_freq = learned_freq
 
 
 
 
121
 
122
- self.register_buffer('dummy', torch.tensor(0), persistent=False)
123
 
124
  self.seq_before_head_dim = seq_before_head_dim
125
  self.default_seq_dim = -3 if seq_before_head_dim else -2
126
 
127
- assert interpolate_factor >= 1.
128
  self.interpolate_factor = interpolate_factor
129
 
130
  self.use_xpos = use_xpos
@@ -135,8 +145,10 @@ class RotaryEmbedding(Module):
135
  scale = (torch.arange(0, dim, 2) + 0.4 * dim) / (1.4 * dim)
136
  self.scale_base = xpos_scale_base
137
 
138
- self.register_buffer('scale', scale, persistent=False)
139
- self.register_buffer('cached_scales', torch.zeros(cache_max_seq_len, dim), persistent=False)
 
 
140
  self.cached_scales_seq_len = 0
141
 
142
  self.apply_rotary_emb = staticmethod(apply_rotary_emb)
@@ -148,11 +160,15 @@ class RotaryEmbedding(Module):
148
  def get_seq_pos(self, seq_len, device=None, dtype=None, offset=0):
149
  device = default(device, self.device)
150
  dtype = default(dtype, self.cached_freqs.dtype)
151
- return (torch.arange(seq_len, device=device, dtype=dtype) + offset) / self.interpolate_factor
 
 
152
 
153
  def rotate_queries_or_keys(self, t, seq_dim=None, offset=0, scale=None):
154
  seq_dim = default(seq_dim, self.default_seq_dim)
155
- assert not self.use_xpos or exists(scale), 'you must use `.rotate_queries_and_keys` method instead'
 
 
156
 
157
  device, dtype, seq_len = t.device, t.dtype, t.shape[seq_dim]
158
  seq = self.get_seq_pos(seq_len, device=device, dtype=dtype, offset=offset)
@@ -161,23 +177,29 @@ class RotaryEmbedding(Module):
161
  if seq_dim == -3:
162
  freqs = freqs.unsqueeze(1)
163
 
164
- return apply_rotary_emb(freqs, t, scale=default(scale, 1.), seq_dim=seq_dim)
165
 
166
  def rotate_queries_with_cached_keys(self, q, k, seq_dim=None, offset=0):
167
- dtype, device, seq_dim = q.dtype, q.device, default(seq_dim, self.default_seq_dim)
 
 
 
 
168
 
169
  q_len, k_len = q.shape[seq_dim], k.shape[seq_dim]
170
  assert q_len <= k_len
171
 
172
- q_scale = k_scale = 1.
173
 
174
  if self.use_xpos:
175
  seq = self.get_seq_pos(k_len, dtype=dtype, device=device)
176
  q_scale = self.get_scale(seq[-q_len:]).type(dtype)
177
  k_scale = self.get_scale(seq).type(dtype)
178
 
179
- rotated_q = self.rotate_queries_or_keys(q, seq_dim=seq_dim, scale=q_scale, offset=k_len - q_len + offset)
180
- rotated_k = self.rotate_queries_or_keys(k, seq_dim=seq_dim, scale=k_scale ** -1)
 
 
181
 
182
  return rotated_q.type(q.dtype), rotated_k.type(k.dtype)
183
 
@@ -195,18 +217,22 @@ class RotaryEmbedding(Module):
195
  scale = scale.unsqueeze(1)
196
 
197
  rotated_q = apply_rotary_emb(freqs, q, scale=scale, seq_dim=seq_dim)
198
- rotated_k = apply_rotary_emb(freqs, k, scale=scale ** -1, seq_dim=seq_dim)
199
 
200
  return rotated_q.type(q.dtype), rotated_k.type(k.dtype)
201
 
202
  def get_scale(self, t: Tensor, seq_len: int | None = None, offset=0):
203
  assert self.use_xpos
204
- should_cache = self.cache_if_possible and exists(seq_len) and (offset + seq_len) <= self.cache_max_seq_len
 
 
 
 
205
 
206
  if should_cache and (seq_len + offset) <= self.cached_scales_seq_len:
207
- return self.cached_scales[offset:(offset + seq_len)]
208
 
209
- scale = 1.
210
  if self.use_xpos:
211
  power = (t - len(t) // 2) / self.scale_base
212
  scale = self.scale ** power.unsqueeze(-1)
@@ -218,7 +244,9 @@ class RotaryEmbedding(Module):
218
 
219
  return scale
220
 
221
- def get_axial_freqs(self, *dims, offsets: tuple[int | float, ...] | Tensor | None = None):
 
 
222
  Colon = slice(None)
223
  all_freqs = []
224
 
@@ -232,7 +260,7 @@ class RotaryEmbedding(Module):
232
  if exists(offsets):
233
  offset = offsets[ind]
234
 
235
- if self.freqs_for == 'pixel':
236
  pos = torch.linspace(-1, 1, steps=dim, device=self.device)
237
  else:
238
  pos = torch.arange(dim, device=self.device)
@@ -248,23 +276,26 @@ class RotaryEmbedding(Module):
248
  all_freqs = broadcast_tensors(*all_freqs)
249
  return torch.cat(all_freqs, dim=-1)
250
 
251
- @autocast('cuda', enabled=False)
252
  def forward(self, t: Tensor, seq_len: int | None = None, offset=0):
253
  should_cache = (
254
- self.cache_if_possible and not self.learned_freq and
255
- exists(seq_len) and self.freqs_for != 'pixel' and
256
- (offset + seq_len) <= self.cache_max_seq_len
 
 
257
  )
258
 
259
  if should_cache and (offset + seq_len) <= self.cached_freqs_seq_len:
260
- return self.cached_freqs[offset:(offset + seq_len)].detach()
261
 
262
  freqs = self.freqs
263
- freqs = torch.einsum('..., f -> ... f', t.type(freqs.dtype), freqs)
264
  freqs = freqs.repeat_interleave(2, dim=-1)
265
 
266
  if should_cache and offset == 0:
267
  self.cached_freqs[:seq_len] = freqs.detach()
268
  self.cached_freqs_seq_len = seq_len
269
 
270
- return freqs
 
 
10
  def exists(val):
11
  return val is not None
12
 
13
+
14
  def default(val, d):
15
  return val if exists(val) else d
16
 
17
+
18
  def broadcat(tensors, dim=-1):
19
  broadcasted_tensors = broadcast_tensors(*tensors)
20
  return torch.cat(broadcasted_tensors, dim=dim)
21
 
22
+
23
  def slice_at_dim(t, dim_slice: slice, *, dim):
24
+ dim += t.ndim if dim < 0 else 0
25
  colons = [slice(None)] * t.ndim
26
  colons[dim] = dim_slice
27
  return t[tuple(colons)]
28
 
29
+
30
  def rotate_half(x):
31
  orig_shape = x.shape
32
  d_head = orig_shape[-1]
33
  x = x.view(*orig_shape[:-1], d_head // 2, 2)
34
+
35
  x1 = x[..., 0]
36
  x2 = x[..., 1]
37
+
38
  res = torch.stack((-x2, x1), dim=-1)
39
  return res.view(*orig_shape)
40
 
41
 
42
+ @autocast("cuda", enabled=False)
43
  def apply_rotary_emb(
44
+ freqs, t, start_index=0, scale=1.0, seq_dim=-2, freqs_seq_dim=None
 
 
 
 
 
45
  ):
46
  dtype = t.dtype
47
 
 
56
  rot_dim = freqs.shape[-1]
57
  end_index = start_index + rot_dim
58
 
59
+ assert rot_dim <= t.shape[-1], (
60
+ f"feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}"
61
+ )
62
 
63
  t_left = t[..., :start_index]
64
  t_middle = t[..., start_index:end_index]
65
  t_right = t[..., end_index:]
66
 
67
+ t_transformed = (t_middle * freqs.cos() * scale) + (
68
+ rotate_half(t_middle) * freqs.sin() * scale
69
+ )
70
+
71
  out = torch.cat((t_left, t_transformed, t_right), dim=-1)
72
  return out.type(dtype)
73
 
74
 
75
  def apply_learned_rotations(rotations, t, start_index=0, freq_ranges=None):
76
  if exists(freq_ranges):
77
+ rotations = torch.einsum("..., f -> ... f", rotations, freq_ranges)
78
  rotations = rotations.reshape(*rotations.shape[:-2], -1)
79
 
80
  rotations = rotations.repeat_interleave(2, dim=-1)
 
86
  self,
87
  dim,
88
  custom_freqs: Tensor | None = None,
89
+ freqs_for: Literal["lang", "pixel", "constant"] = "lang",
90
+ theta=10000,
91
+ max_freq=10,
92
+ num_freqs=1,
93
+ learned_freq=False,
94
+ use_xpos=False,
95
+ xpos_scale_base=512,
96
+ interpolate_factor=1.0,
97
+ theta_rescale_factor=1.0,
98
+ seq_before_head_dim=False,
99
+ cache_if_possible=True,
100
+ cache_max_seq_len=8192,
101
  ):
102
  super().__init__()
103
 
 
106
 
107
  if exists(custom_freqs):
108
  freqs = custom_freqs
109
+ elif freqs_for == "lang":
110
+ freqs = 1.0 / (
111
+ theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
112
+ )
113
+ elif freqs_for == "pixel":
114
+ freqs = torch.linspace(1.0, max_freq / 2, dim // 2) * pi
115
+ elif freqs_for == "constant":
116
  freqs = torch.ones(num_freqs).float()
117
 
118
  self.cache_if_possible = cache_if_possible
119
  self.cache_max_seq_len = cache_max_seq_len
120
 
121
+ self.register_buffer(
122
+ "cached_freqs", torch.zeros(cache_max_seq_len, dim), persistent=False
123
+ )
124
  self.cached_freqs_seq_len = 0
125
 
 
126
  self.learned_freq = learned_freq
127
+ if learned_freq:
128
+ self.freqs = nn.Parameter(freqs)
129
+ else:
130
+ self.register_buffer("freqs", freqs, persistent=False)
131
 
132
+ self.register_buffer("dummy", torch.tensor(0), persistent=False)
133
 
134
  self.seq_before_head_dim = seq_before_head_dim
135
  self.default_seq_dim = -3 if seq_before_head_dim else -2
136
 
137
+ assert interpolate_factor >= 1.0
138
  self.interpolate_factor = interpolate_factor
139
 
140
  self.use_xpos = use_xpos
 
145
  scale = (torch.arange(0, dim, 2) + 0.4 * dim) / (1.4 * dim)
146
  self.scale_base = xpos_scale_base
147
 
148
+ self.register_buffer("scale", scale, persistent=False)
149
+ self.register_buffer(
150
+ "cached_scales", torch.zeros(cache_max_seq_len, dim), persistent=False
151
+ )
152
  self.cached_scales_seq_len = 0
153
 
154
  self.apply_rotary_emb = staticmethod(apply_rotary_emb)
 
160
  def get_seq_pos(self, seq_len, device=None, dtype=None, offset=0):
161
  device = default(device, self.device)
162
  dtype = default(dtype, self.cached_freqs.dtype)
163
+ return (
164
+ torch.arange(seq_len, device=device, dtype=dtype) + offset
165
+ ) / self.interpolate_factor
166
 
167
  def rotate_queries_or_keys(self, t, seq_dim=None, offset=0, scale=None):
168
  seq_dim = default(seq_dim, self.default_seq_dim)
169
+ assert not self.use_xpos or exists(scale), (
170
+ "you must use `.rotate_queries_and_keys` method instead"
171
+ )
172
 
173
  device, dtype, seq_len = t.device, t.dtype, t.shape[seq_dim]
174
  seq = self.get_seq_pos(seq_len, device=device, dtype=dtype, offset=offset)
 
177
  if seq_dim == -3:
178
  freqs = freqs.unsqueeze(1)
179
 
180
+ return apply_rotary_emb(freqs, t, scale=default(scale, 1.0), seq_dim=seq_dim)
181
 
182
  def rotate_queries_with_cached_keys(self, q, k, seq_dim=None, offset=0):
183
+ dtype, device, seq_dim = (
184
+ q.dtype,
185
+ q.device,
186
+ default(seq_dim, self.default_seq_dim),
187
+ )
188
 
189
  q_len, k_len = q.shape[seq_dim], k.shape[seq_dim]
190
  assert q_len <= k_len
191
 
192
+ q_scale = k_scale = 1.0
193
 
194
  if self.use_xpos:
195
  seq = self.get_seq_pos(k_len, dtype=dtype, device=device)
196
  q_scale = self.get_scale(seq[-q_len:]).type(dtype)
197
  k_scale = self.get_scale(seq).type(dtype)
198
 
199
+ rotated_q = self.rotate_queries_or_keys(
200
+ q, seq_dim=seq_dim, scale=q_scale, offset=k_len - q_len + offset
201
+ )
202
+ rotated_k = self.rotate_queries_or_keys(k, seq_dim=seq_dim, scale=k_scale**-1)
203
 
204
  return rotated_q.type(q.dtype), rotated_k.type(k.dtype)
205
 
 
217
  scale = scale.unsqueeze(1)
218
 
219
  rotated_q = apply_rotary_emb(freqs, q, scale=scale, seq_dim=seq_dim)
220
+ rotated_k = apply_rotary_emb(freqs, k, scale=scale**-1, seq_dim=seq_dim)
221
 
222
  return rotated_q.type(q.dtype), rotated_k.type(k.dtype)
223
 
224
  def get_scale(self, t: Tensor, seq_len: int | None = None, offset=0):
225
  assert self.use_xpos
226
+ should_cache = (
227
+ self.cache_if_possible
228
+ and exists(seq_len)
229
+ and (offset + seq_len) <= self.cache_max_seq_len
230
+ )
231
 
232
  if should_cache and (seq_len + offset) <= self.cached_scales_seq_len:
233
+ return self.cached_scales[offset : (offset + seq_len)]
234
 
235
+ scale = 1.0
236
  if self.use_xpos:
237
  power = (t - len(t) // 2) / self.scale_base
238
  scale = self.scale ** power.unsqueeze(-1)
 
244
 
245
  return scale
246
 
247
+ def get_axial_freqs(
248
+ self, *dims, offsets: tuple[int | float, ...] | Tensor | None = None
249
+ ):
250
  Colon = slice(None)
251
  all_freqs = []
252
 
 
260
  if exists(offsets):
261
  offset = offsets[ind]
262
 
263
+ if self.freqs_for == "pixel":
264
  pos = torch.linspace(-1, 1, steps=dim, device=self.device)
265
  else:
266
  pos = torch.arange(dim, device=self.device)
 
276
  all_freqs = broadcast_tensors(*all_freqs)
277
  return torch.cat(all_freqs, dim=-1)
278
 
279
+ @autocast("cuda", enabled=False)
280
  def forward(self, t: Tensor, seq_len: int | None = None, offset=0):
281
  should_cache = (
282
+ self.cache_if_possible
283
+ and not self.learned_freq
284
+ and exists(seq_len)
285
+ and self.freqs_for != "pixel"
286
+ and (offset + seq_len) <= self.cache_max_seq_len
287
  )
288
 
289
  if should_cache and (offset + seq_len) <= self.cached_freqs_seq_len:
290
+ return self.cached_freqs[offset : (offset + seq_len)].detach()
291
 
292
  freqs = self.freqs
293
+ freqs = torch.einsum("..., f -> ... f", t.type(freqs.dtype), freqs)
294
  freqs = freqs.repeat_interleave(2, dim=-1)
295
 
296
  if should_cache and offset == 0:
297
  self.cached_freqs[:seq_len] = freqs.detach()
298
  self.cached_freqs_seq_len = seq_len
299
 
300
+ return freqs
301
+