vector-quantize-pytorch 1.25.2__tar.gz → 1.27.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (34) hide show
  1. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/PKG-INFO +2 -2
  2. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/pyproject.toml +2 -2
  3. vector_quantize_pytorch-1.27.0/tests/test_beam.py +64 -0
  4. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/tests/test_readme.py +4 -2
  5. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/finite_scalar_quantization.py +20 -7
  6. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_fsq.py +22 -11
  7. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_vq.py +180 -27
  8. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/vector_quantize_pytorch.py +176 -278
  9. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.github/workflows/build.yml +0 -0
  10. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.github/workflows/python-publish.yml +0 -0
  11. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.github/workflows/test.yml +0 -0
  12. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.gitignore +0 -0
  13. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/LICENSE +0 -0
  14. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/README.md +0 -0
  15. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder.py +0 -0
  16. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_fsq.py +0 -0
  17. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_lfq.py +0 -0
  18. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_sim_vq.py +0 -0
  19. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/fsq.png +0 -0
  20. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/lfq.png +0 -0
  21. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/simvq.png +0 -0
  22. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/vq.png +0 -0
  23. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/ruff.toml +0 -0
  24. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/tests/test_latent_quantization.py +0 -0
  25. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/tests/test_lfq.py +0 -0
  26. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/__init__.py +0 -0
  27. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/binary_mapper.py +0 -0
  28. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/latent_quantization.py +0 -0
  29. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  30. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  31. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_lfq.py +0 -0
  32. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  33. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/sim_vq.py +0 -0
  34. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.25.2
3
+ Version: 1.27.0
4
4
  Summary: Vector Quantization - Pytorch
5
5
  Project-URL: Homepage, https://pypi.org/project/vector-quantize-pytorch/
6
6
  Project-URL: Repository, https://github.com/lucidrains/vector-quantizer-pytorch
@@ -36,7 +36,7 @@ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
36
36
  Requires-Python: >=3.9
37
37
  Requires-Dist: einops>=0.8.0
38
38
  Requires-Dist: einx>=0.3.0
39
- Requires-Dist: torch>=2.0
39
+ Requires-Dist: torch>=2.4
40
40
  Provides-Extra: examples
41
41
  Requires-Dist: torchvision; extra == 'examples'
42
42
  Requires-Dist: tqdm; extra == 'examples'
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.25.2"
3
+ version = "1.27.0"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -23,7 +23,7 @@ classifiers=[
23
23
  ]
24
24
 
25
25
  dependencies = [
26
- "torch>=2.0",
26
+ "torch>=2.4",
27
27
  "einops>=0.8.0",
28
28
  "einx>=0.3.0",
29
29
  ]
@@ -0,0 +1,64 @@
1
+ import torch
2
+ from vector_quantize_pytorch import VectorQuantize
3
+
4
+ def test_topk_and_manual_ema_update():
5
+
6
+ vq1 = VectorQuantize(
7
+ dim = 256,
8
+ codebook_size = 512
9
+ )
10
+
11
+ vq2 = VectorQuantize(
12
+ dim = 256,
13
+ codebook_size = 512
14
+ )
15
+
16
+ vq2.load_state_dict(vq1.state_dict())
17
+
18
+ x = torch.randn(1, 1024, 256)
19
+ mask = torch.randint(0, 2, (1, 1024)).bool()
20
+
21
+ vq1.train()
22
+ quantize1, indices1, commit_loss1 = vq1(x, mask = mask)
23
+
24
+ vq2.train()
25
+ quantize2, indices2, commit_losses = vq2(x, mask = mask, topk = 1, ema_update = False)
26
+
27
+ assert quantize2.shape == (1, 1024, 1, 256)
28
+ assert indices2.shape == (1, 1024, 1)
29
+ assert commit_losses.shape == (1, 1024, 1)
30
+
31
+ top_quantize2 = quantize2[..., 0, :]
32
+ top_indices2 = indices2[..., 0]
33
+
34
+ assert torch.allclose(commit_loss1, commit_losses.sum() / mask.sum())
35
+ assert torch.equal(indices1, top_indices2)
36
+ assert torch.allclose(quantize1, top_quantize2)
37
+
38
+ assert not torch.allclose(vq1._codebook.embed_avg, vq2._codebook.embed_avg)
39
+
40
+ vq2.update_ema_indices(x, top_indices2, mask = mask)
41
+
42
+ assert torch.allclose(vq1._codebook.cluster_size, vq2._codebook.cluster_size)
43
+ assert torch.allclose(vq1._codebook.embed_avg, vq2._codebook.embed_avg)
44
+ assert torch.allclose(vq1.codebook, vq2.codebook)
45
+
46
+ def test_beam_search():
47
+ import torch
48
+ from vector_quantize_pytorch import ResidualVQ
49
+
50
+ residual_vq = ResidualVQ(
51
+ dim = 256,
52
+ num_quantizers = 8, # specify number of quantizers
53
+ codebook_size = 1024, # codebook size
54
+ quantize_dropout = True
55
+ )
56
+
57
+ x = torch.randn(1, 1024, 256)
58
+
59
+ for _ in range(5):
60
+ quantized, indices, commit_loss = residual_vq(x, beam_size = 3)
61
+
62
+ assert quantized.shape == (1, 1024, 256)
63
+ assert indices.shape == (1, 1024, 8)
64
+ assert commit_loss.shape == (8,)
@@ -247,13 +247,15 @@ def test_directional_reparam():
247
247
  quantized, indices, _ = rq(x)
248
248
 
249
249
  @pytest.mark.parametrize('preserve_symmetry', (True, False))
250
+ @pytest.mark.parametrize('bound_hard_clamp', (True, False))
250
251
  def test_fsq(
251
- preserve_symmetry
252
+ preserve_symmetry,
253
+ bound_hard_clamp
252
254
  ):
253
255
  from vector_quantize_pytorch import FSQ
254
256
 
255
257
  levels = [8,5,5,5] # see 4.1 and A.4.1 in the paper
256
- quantizer = FSQ(levels, preserve_symmetry = preserve_symmetry)
258
+ quantizer = FSQ(levels, preserve_symmetry = preserve_symmetry, bound_hard_clamp = bound_hard_clamp)
257
259
 
258
260
  x = torch.randn(1, 1024, 4) # 4 since there are 4 levels
259
261
  xhat, indices = quantizer(x)
@@ -11,7 +11,7 @@ from typing import List, Tuple
11
11
  import torch
12
12
  import torch.nn as nn
13
13
  from torch.nn import Module
14
- from torch import tensor, Tensor, int32
14
+ from torch import tensor, Tensor, int32, tanh, atanh, clamp
15
15
  from torch.amp import autocast
16
16
 
17
17
  import einx
@@ -30,6 +30,9 @@ def default(*args):
30
30
  return arg
31
31
  return None
32
32
 
33
+ def identity(t):
34
+ return t
35
+
33
36
  def maybe(fn):
34
37
  @wraps(fn)
35
38
  def inner(x, *args, **kwargs):
@@ -73,6 +76,7 @@ class FSQ(Module):
73
76
  force_quantization_f32 = True,
74
77
  preserve_symmetry = False,
75
78
  noise_dropout = 0.,
79
+ bound_hard_clamp = False # for residual fsq, if input is pre-softclamped to the right range
76
80
  ):
77
81
  super().__init__()
78
82
 
@@ -121,22 +125,31 @@ class FSQ(Module):
121
125
  self.allowed_dtypes = allowed_dtypes
122
126
  self.force_quantization_f32 = force_quantization_f32
123
127
 
124
- def bound(self, z, eps = 1e-3):
128
+ # allow for a hard clamp
129
+
130
+ self.bound_hard_clamp = bound_hard_clamp
131
+
132
+ def bound(self, z, eps = 1e-3, hard_clamp = False):
125
133
  """ Bound `z`, an array of shape (..., d). """
134
+ maybe_tanh = tanh if not hard_clamp else partial(clamp, min = -1., max = 1.)
135
+ maybe_atanh = atanh if not hard_clamp else identity
136
+
126
137
  half_l = (self._levels - 1) * (1 + eps) / 2
127
138
  offset = torch.where(self._levels % 2 == 0, 0.5, 0.0)
128
- shift = (offset / half_l).atanh()
129
- bounded_z = (z + shift).tanh() * half_l - offset
139
+ shift = maybe_atanh(offset / half_l)
140
+ bounded_z = maybe_tanh(z + shift) * half_l - offset
130
141
  half_width = self._levels // 2
131
142
  return round_ste(bounded_z) / half_width
132
143
 
133
144
  # symmetry-preserving and noise-approximated quantization, section 3.2 in https://arxiv.org/abs/2411.19842
134
145
 
135
- def symmetry_preserving_bound(self, z):
146
+ def symmetry_preserving_bound(self, z, hard_clamp = False):
136
147
  """ QL(x) = 2 / (L - 1) * [(L - 1) * (tanh(x) + 1) / 2 + 0.5] - 1 """
148
+ maybe_tanh = tanh if not hard_clamp else partial(clamp, min = -1., max = 1.)
149
+
137
150
  levels_minus_1 = (self._levels - 1)
138
151
  scale = 2. / levels_minus_1
139
- bracket = (levels_minus_1 * (z.tanh() + 1) / 2.) + 0.5
152
+ bracket = (levels_minus_1 * (maybe_tanh(z) + 1) / 2.) + 0.5
140
153
  bracket = floor_ste(bracket)
141
154
  return scale * bracket - 1.
142
155
 
@@ -146,7 +159,7 @@ class FSQ(Module):
146
159
  shape, device, noise_dropout, preserve_symmetry = z.shape[0], z.device, self.noise_dropout, self.preserve_symmetry
147
160
  bound_fn = self.symmetry_preserving_bound if preserve_symmetry else self.bound
148
161
 
149
- bounded_z = bound_fn(z)
162
+ bounded_z = bound_fn(z, hard_clamp = self.bound_hard_clamp)
150
163
 
151
164
  # determine where to add a random offset elementwise
152
165
  # if using noise dropout
@@ -1,11 +1,11 @@
1
+ from __future__ import annotations
2
+
1
3
  import random
2
4
  from math import ceil
3
5
  from functools import partial
4
6
 
5
- from typing import List
6
-
7
7
  import torch
8
- from torch import nn
8
+ from torch import nn, tensor
9
9
  from torch.nn import Module, ModuleList
10
10
  import torch.nn.functional as F
11
11
  from torch.amp import autocast
@@ -52,14 +52,15 @@ class ResidualFSQ(Module):
52
52
  def __init__(
53
53
  self,
54
54
  *,
55
- levels: List[int],
55
+ levels: list[int],
56
56
  num_quantizers,
57
57
  dim = None,
58
58
  is_channel_first = False,
59
59
  quantize_dropout = False,
60
60
  quantize_dropout_cutoff_index = 0,
61
61
  quantize_dropout_multiple_of = 1,
62
- soft_clamp_input_value = None,
62
+ soft_clamp_input_value: float | list[float] | Tensor | None = None,
63
+ bound_hard_clamp = True,
63
64
  **kwargs
64
65
  ):
65
66
  super().__init__()
@@ -74,25 +75,24 @@ class ResidualFSQ(Module):
74
75
  self.is_channel_first = is_channel_first
75
76
  self.num_quantizers = num_quantizers
76
77
 
77
- # soft clamping the input value
78
-
79
- self.soft_clamp_input_value = soft_clamp_input_value
80
-
81
78
  # layers
82
79
 
83
80
  self.levels = levels
84
81
  self.layers = nn.ModuleList([])
85
82
 
86
- levels_tensor = torch.Tensor(levels)
83
+ levels_tensor = tensor(levels)
84
+ assert (levels_tensor > 1).all()
87
85
 
88
86
  scales = []
89
87
 
90
88
  for ind in range(num_quantizers):
91
- scales.append((levels_tensor - 1) ** -ind)
89
+ scales.append(levels_tensor.float() ** -ind)
92
90
 
93
91
  fsq = FSQ(
94
92
  levels = levels,
95
93
  dim = codebook_dim,
94
+ preserve_symmetry = True,
95
+ bound_hard_clamp = bound_hard_clamp,
96
96
  **kwargs
97
97
  )
98
98
 
@@ -111,6 +111,17 @@ class ResidualFSQ(Module):
111
111
  self.quantize_dropout_cutoff_index = quantize_dropout_cutoff_index
112
112
  self.quantize_dropout_multiple_of = quantize_dropout_multiple_of # encodec paper proposes structured dropout, believe this was set to 4
113
113
 
114
+ # soft clamping the input value
115
+
116
+ if bound_hard_clamp:
117
+ assert not exists(soft_clamp_input_value)
118
+ soft_clamp_input_value = 1 + (1 / (levels_tensor - 1))
119
+
120
+ if isinstance(soft_clamp_input_value, (list, float)):
121
+ soft_clamp_input_value = tensor(soft_clamp_input_value)
122
+
123
+ self.register_buffer('soft_clamp_input_value', soft_clamp_input_value, persistent = False)
124
+
114
125
  @property
115
126
  def codebooks(self):
116
127
  codebooks = [layer.implicit_codebook for layer in self.layers]
@@ -6,7 +6,7 @@ from functools import partial, cache
6
6
  from itertools import zip_longest
7
7
 
8
8
  import torch
9
- from torch import nn, Tensor
9
+ from torch import nn, Tensor, arange, cat
10
10
  from torch.nn import Module, ModuleList
11
11
  import torch.nn.functional as F
12
12
  import torch.distributed as dist
@@ -14,6 +14,7 @@ from vector_quantize_pytorch.vector_quantize_pytorch import VectorQuantize
14
14
 
15
15
  from einops import rearrange, repeat, reduce, pack, unpack
16
16
 
17
+ import einx
17
18
  from einx import get_at
18
19
 
19
20
  # helper functions
@@ -36,6 +37,47 @@ def unique(arr):
36
37
  def round_up_multiple(num, mult):
37
38
  return ceil(num / mult) * mult
38
39
 
40
+ # tensor helpers
41
+
42
+ def pad_at_dim(
43
+ t,
44
+ pad: tuple[int, int],
45
+ dim = -1,
46
+ value = 0.
47
+ ):
48
+ if pad == (0, 0):
49
+ return t
50
+
51
+ dims_from_right = (- dim - 1) if dim < 0 else (t.ndim - dim - 1)
52
+ zeros = ((0, 0) * dims_from_right)
53
+ return F.pad(t, (*zeros, *pad), value = value)
54
+
55
+ def pack_one(t, pattern):
56
+ packed, packed_shape = pack([t], pattern)
57
+
58
+ def inverse(out, inv_pattern = None):
59
+ inv_pattern = default(inv_pattern, pattern)
60
+ return first(unpack(out, packed_shape, inv_pattern))
61
+
62
+ return packed, inverse
63
+
64
+ def batch_select(t, indices, pattern = None):
65
+
66
+ if exists(pattern):
67
+ indices = rearrange(indices, '... k -> (...) k')
68
+ t, inv_pack = pack_one(t, pattern)
69
+
70
+
71
+ batch_indices = arange(t.shape[0], device = t.device)
72
+ batch_indices = rearrange(batch_indices, 'b -> b 1')
73
+
74
+ out = t[batch_indices, indices]
75
+
76
+ if exists(pattern):
77
+ out = inv_pack(out)
78
+
79
+ return out
80
+
39
81
  # distributed helpers
40
82
 
41
83
  def is_distributed():
@@ -128,6 +170,8 @@ class ResidualVQ(Module):
128
170
  accept_image_fmap = False,
129
171
  implicit_neural_codebook = False, # QINCo from https://arxiv.org/abs/2401.14732
130
172
  mlp_kwargs: dict = dict(),
173
+ beam_size = None,
174
+ eval_beam_size = None,
131
175
  **vq_kwargs
132
176
  ):
133
177
  super().__init__()
@@ -183,6 +227,15 @@ class ResidualVQ(Module):
183
227
  self.quantize_dropout_cutoff_index = quantize_dropout_cutoff_index
184
228
  self.quantize_dropout_multiple_of = quantize_dropout_multiple_of # encodec paper proposes structured dropout, believe this was set to 4
185
229
 
230
+ # determine whether is using ema update
231
+
232
+ self.vq_is_ema_updating = first(self.layers).ema_update
233
+
234
+ # beam size
235
+
236
+ self.beam_size = default(beam_size, eval_beam_size)
237
+ self.eval_beam_size = eval_beam_size
238
+
186
239
  # setting up the MLPs for implicit neural codebooks
187
240
 
188
241
  self.mlps = None
@@ -295,19 +348,29 @@ class ResidualVQ(Module):
295
348
  return_all_codes = False,
296
349
  sample_codebook_temp = None,
297
350
  freeze_codebook = False,
351
+ beam_size = None,
298
352
  rand_quantize_dropout_fixed_seed = None
299
353
  ):
300
- num_quant, quant_dropout_multiple_of, return_loss, device = self.num_quantizers, self.quantize_dropout_multiple_of, exists(indices), x.device
354
+
355
+ # variables
356
+
357
+ input_shape, num_quant, quant_dropout_multiple_of, return_loss, device = x.shape, self.num_quantizers, self.quantize_dropout_multiple_of, exists(indices), x.device
358
+
359
+ beam_size = default(beam_size, self.beam_size if self.training else self.eval_beam_size)
360
+
361
+ is_beam_search = exists(beam_size) and beam_size > 1
362
+
363
+ # projecting in
301
364
 
302
365
  x = self.project_in(x)
303
366
 
304
367
  assert not (self.accept_image_fmap and exists(indices))
305
368
 
306
- quantized_out = 0.
369
+ quantized_out = torch.zeros_like(x)
307
370
  residual = x
308
371
 
309
- all_losses = []
310
- all_indices = []
372
+ all_losses = torch.empty((0,), device = device, dtype = torch.float32)
373
+ all_indices = torch.empty((*input_shape[:-1], 0), device = device, dtype = torch.long)
311
374
 
312
375
  if isinstance(indices, list):
313
376
  indices = torch.stack(indices)
@@ -319,7 +382,6 @@ class ResidualVQ(Module):
319
382
  should_quantize_dropout = self.training and self.quantize_dropout and not return_loss
320
383
 
321
384
  # sample a layer index at which to dropout further residual quantization
322
- # also prepare null indices and loss
323
385
 
324
386
  if should_quantize_dropout:
325
387
 
@@ -335,9 +397,24 @@ class ResidualVQ(Module):
335
397
  if quant_dropout_multiple_of != 1:
336
398
  rand_quantize_dropout_index = round_up_multiple(rand_quantize_dropout_index + 1, quant_dropout_multiple_of) - 1
337
399
 
338
- null_indices_shape = (x.shape[0], *x.shape[-2:]) if self.accept_image_fmap else tuple(x.shape[:2])
339
- null_indices = torch.full(null_indices_shape, -1., device = device, dtype = torch.long)
340
- null_loss = torch.full((1,), 0., device = device, dtype = x.dtype)
400
+ # save all inputs across layers, for use during expiration at end under shared codebook setting, or ema update during beam search
401
+
402
+ all_residuals = torch.empty((*input_shape[:-1], 0, input_shape[-1]), dtype = residual.dtype, device = device)
403
+
404
+ # maybe prepare beam search
405
+
406
+ if is_beam_search:
407
+ prec_dims = x.shape[:-1]
408
+
409
+ search_scores = torch.zeros((*prec_dims, 1), device = device, dtype = x.dtype)
410
+
411
+ residual = rearrange(residual, '... d -> ... 1 d')
412
+ quantized_out = rearrange(quantized_out, '... d -> ... 1 d')
413
+
414
+ all_residuals = rearrange(all_residuals, '... l d -> ... 1 l d')
415
+ all_indices = rearrange(all_indices, '... l -> ... 1 l')
416
+
417
+ all_losses = all_losses.reshape(*input_shape[:-1], 1, 0)
341
418
 
342
419
  # setup the mlps for implicit neural codebook
343
420
 
@@ -346,17 +423,13 @@ class ResidualVQ(Module):
346
423
  if self.implicit_neural_codebook:
347
424
  maybe_code_transforms = (None, *self.mlps)
348
425
 
349
- # save all inputs across layers, for use during expiration at end under shared codebook setting
350
-
351
- all_residuals = []
352
-
353
426
  # go through the layers
354
427
 
355
428
  for quantizer_index, (vq, maybe_mlp) in enumerate(zip(self.layers, maybe_code_transforms)):
356
429
 
357
430
  if should_quantize_dropout and quantizer_index > rand_quantize_dropout_index:
358
- all_indices.append(null_indices)
359
- all_losses.append(null_loss)
431
+ all_indices = pad_at_dim(all_indices, (0, 1), value = -1, dim = -1)
432
+ all_losses = pad_at_dim(all_losses, (0, 1), value = 0, dim = -1)
360
433
  continue
361
434
 
362
435
  layer_indices = None
@@ -368,9 +441,9 @@ class ResidualVQ(Module):
368
441
  if exists(maybe_mlp):
369
442
  maybe_mlp = partial(maybe_mlp, condition = quantized_out)
370
443
 
371
- # save for expiration
444
+ # save the residual input for maybe expiration as well as ema update after beam search
372
445
 
373
- all_residuals.append(residual)
446
+ all_residuals = cat((all_residuals, rearrange(residual, '... d -> ... 1 d')), dim = -2)
374
447
 
375
448
  # vector quantize forward
376
449
 
@@ -380,29 +453,111 @@ class ResidualVQ(Module):
380
453
  indices = layer_indices,
381
454
  sample_codebook_temp = sample_codebook_temp,
382
455
  freeze_codebook = freeze_codebook,
383
- codebook_transform_fn = maybe_mlp
456
+ codebook_transform_fn = maybe_mlp,
457
+ topk = beam_size if is_beam_search else None
384
458
  )
385
459
 
386
- residual = residual - quantized.detach()
387
- quantized_out = quantized_out + quantized
460
+ # cross entropy loss for some old paper
388
461
 
389
462
  if return_loss:
390
- ce_loss = rest[0]
463
+ ce_loss = first(rest)
391
464
  ce_losses.append(ce_loss)
392
465
  continue
393
466
 
394
467
  embed_indices, loss = rest
395
468
 
396
- all_indices.append(embed_indices)
397
- all_losses.append(loss)
469
+ # handle expanding first residual if doing beam search
470
+
471
+ if is_beam_search:
472
+
473
+ search_scores = einx.add('... j, ... j k -> ... (j k)', search_scores, -loss)
474
+
475
+ residual = rearrange(residual, '... j d -> ... j 1 d')
476
+ quantized_out = rearrange(quantized_out, '... j d -> ... j 1 d')
477
+
478
+ all_residuals = repeat(all_residuals, '... j l d -> ... (j k) l d', k = beam_size)
479
+
480
+ # core residual vq logic
481
+
482
+ residual = residual - quantized.detach()
483
+ quantized_out = quantized_out + quantized
484
+
485
+ # handle sort and topk beams
486
+
487
+ if is_beam_search:
488
+ residual = rearrange(residual, '... j k d -> ... (j k) d')
489
+ quantized_out = rearrange(quantized_out, '... j k d -> ... (j k) d')
490
+
491
+ # broadcat the indices
492
+
493
+ all_indices = repeat(all_indices, '... j l -> ... j k l', k = embed_indices.shape[-1])
494
+ embed_indices = rearrange(embed_indices, '... j k -> ... j k 1')
495
+
496
+ all_indices = cat((all_indices, embed_indices), dim = -1)
497
+ all_indices = rearrange(all_indices, '... j k l -> ... (j k) l')
498
+
499
+ # broadcat the losses
500
+
501
+ all_losses = repeat(all_losses, '... j l -> ... j k l', k = loss.shape[-1])
502
+ loss = rearrange(loss, '... -> ... 1')
503
+
504
+ all_losses = cat((all_losses, loss), dim = -1)
505
+ all_losses = rearrange(all_losses, '... j k l -> ... (j k) l')
506
+
507
+ # handle sort and selection of highest beam size
508
+
509
+ if search_scores.shape[-1] > beam_size:
510
+ search_scores, select_indices = search_scores.topk(beam_size, dim = -1)
511
+
512
+ residual = batch_select(residual, select_indices, '* k d')
513
+ quantized_out = batch_select(quantized_out, select_indices, '* k d')
514
+
515
+ all_indices = batch_select(all_indices, select_indices, '* k l')
516
+ all_losses = batch_select(all_losses, select_indices, '* k l')
517
+
518
+ all_residuals = batch_select(all_residuals, select_indices, '* k l d')
519
+ else:
520
+ # aggregate indices and losses
521
+
522
+ all_indices = cat((all_indices, rearrange(embed_indices, '... -> ... 1')), dim = -1)
523
+
524
+ all_losses = cat((all_losses, rearrange(loss, '... -> ... 1')), dim = -1)
525
+
526
+ # handle beam search
527
+
528
+ if is_beam_search:
529
+ top_index = search_scores.argmax(dim = -1, keepdim = True)
530
+
531
+ quantized_out = batch_select(quantized_out, top_index, '* k d')
532
+ all_indices = batch_select(all_indices, top_index, '* k l')
533
+ all_losses = batch_select(all_losses, top_index, '* k l')
534
+ all_residuals = batch_select(all_residuals, top_index, '* k l d')
535
+
536
+ quantized_out, all_indices, all_losses, all_residuals = [t[..., 0, :] for t in (quantized_out, all_indices, all_losses, all_residuals)]
537
+
538
+ # handle commit loss, which should be the average
539
+
540
+ if exists(mask):
541
+ all_losses = einx.where('..., ... l,', mask, all_losses, 0.)
542
+ all_losses = reduce(all_losses, '... l -> l', 'sum') / mask.sum(dim = -1).clamp_min(1e-4)
543
+ else:
544
+ all_losses = reduce(all_losses, '... l -> l', 'mean')
545
+
546
+ # handle updating ema
547
+
548
+ if self.vq_is_ema_updating:
549
+ for vq, layer_input, indices in zip(self.layers, all_residuals.unbind(dim = -2), all_indices.unbind(dim = -1)): # in the case of quantize dropout, zip will terminate with the shorter sequence, which should be all_residuals
550
+ vq.update_ema_indices(layer_input, indices, mask = mask)
398
551
 
399
552
  # if shared codebook, update ema only at end
400
553
 
401
- if self.training and self.shared_codebook:
554
+ if self.training and self.shared_codebook and not is_beam_search:
402
555
  shared_layer = first(self.layers)
403
556
  shared_layer._codebook.update_ema()
404
557
  shared_layer.update_in_place_optimizer()
405
- shared_layer.expire_codes_(torch.cat(all_residuals, dim = -2))
558
+
559
+ all_codes_for_expire = rearrange(all_residuals, '... n l d -> ... (n l) d')
560
+ shared_layer.expire_codes_(all_codes_for_expire)
406
561
 
407
562
  # project out, if needed
408
563
 
@@ -415,8 +570,6 @@ class ResidualVQ(Module):
415
570
 
416
571
  # stack all losses and indices
417
572
 
418
- all_losses, all_indices = map(partial(torch.stack, dim = -1), (all_losses, all_indices))
419
-
420
573
  ret = (quantized_out, all_indices, all_losses)
421
574
 
422
575
  if return_all_codes: