vector-quantize-pytorch 1.26.0__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.26.0 → vector_quantize_pytorch-1.27.0}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/pyproject.toml +1 -1
  3. vector_quantize_pytorch-1.27.0/tests/test_beam.py +64 -0
  4. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_vq.py +180 -27
  5. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/vector_quantize_pytorch.py +176 -278
  6. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/build.yml +0 -0
  7. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/python-publish.yml +0 -0
  8. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/test.yml +0 -0
  9. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.gitignore +0 -0
  10. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/LICENSE +0 -0
  11. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/README.md +0 -0
  12. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder.py +0 -0
  13. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_fsq.py +0 -0
  14. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_lfq.py +0 -0
  15. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_sim_vq.py +0 -0
  16. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/fsq.png +0 -0
  17. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/lfq.png +0 -0
  18. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/simvq.png +0 -0
  19. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/vq.png +0 -0
  20. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/ruff.toml +0 -0
  21. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/tests/test_latent_quantization.py +0 -0
  22. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/tests/test_lfq.py +0 -0
  23. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/tests/test_readme.py +0 -0
  24. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/__init__.py +0 -0
  25. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/binary_mapper.py +0 -0
  26. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  27. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/latent_quantization.py +0 -0
  28. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  29. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  30. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_fsq.py +0 -0
  31. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_lfq.py +0 -0
  32. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  33. {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/sim_vq.py +0 -0
  34. {vector_quantize_pytorch-1.26.0 → 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.26.0
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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.26.0"
3
+ version = "1.27.0"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -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,)
@@ -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:
@@ -1,4 +1,5 @@
1
1
  from __future__ import annotations
2
+ from typing import Callable
2
3
 
3
4
  from math import sqrt
4
5
  from functools import partial, cache
@@ -6,7 +7,7 @@ from collections import namedtuple
6
7
 
7
8
  import torch
8
9
  from torch.nn import Module
9
- from torch import nn, einsum, is_tensor, Tensor
10
+ from torch import nn, einsum, tensor, is_tensor, Tensor
10
11
  import torch.nn.functional as F
11
12
  import torch.distributed as distributed
12
13
  from torch.optim import Optimizer
@@ -15,8 +16,6 @@ from torch.amp import autocast
15
16
  import einx
16
17
  from einops import rearrange, repeat, reduce, pack, unpack
17
18
 
18
- from typing import Callable
19
-
20
19
  def exists(val):
21
20
  return val is not None
22
21
 
@@ -120,7 +119,8 @@ def gumbel_sample(
120
119
  stochastic = False,
121
120
  straight_through = False,
122
121
  dim = -1,
123
- training = True
122
+ training = True,
123
+ topk = None
124
124
  ):
125
125
  dtype, size = logits.dtype, logits.shape[dim]
126
126
 
@@ -129,7 +129,11 @@ def gumbel_sample(
129
129
  else:
130
130
  sampling_logits = logits
131
131
 
132
- ind = sampling_logits.argmax(dim = dim)
132
+ if exists(topk):
133
+ ind = sampling_logits.topk(topk, dim = dim).indices
134
+ else:
135
+ ind = sampling_logits.argmax(dim = dim)
136
+
133
137
  one_hot = F.one_hot(ind, size).type(dtype)
134
138
 
135
139
  if not straight_through or temperature <= 0. or not training:
@@ -182,7 +186,7 @@ def sample_multinomial(total_count, probs):
182
186
  return sample.to(device)
183
187
 
184
188
  def all_gather_sizes(x, dim):
185
- size = torch.tensor(x.shape[dim], dtype = torch.long, device = x.device)
189
+ size = tensor(x.shape[dim], dtype = torch.long, device = x.device)
186
190
  all_sizes = [torch.empty_like(size) for _ in range(distributed.get_world_size())]
187
191
  distributed.all_gather(all_sizes, size)
188
192
  return torch.stack(all_sizes)
@@ -337,7 +341,7 @@ def orthogonal_loss_fn(t):
337
341
 
338
342
  # distance types
339
343
 
340
- class EuclideanCodebook(Module):
344
+ class Codebook(Module):
341
345
  def __init__(
342
346
  self,
343
347
  dim,
@@ -359,17 +363,23 @@ class EuclideanCodebook(Module):
359
363
  affine_param = False,
360
364
  sync_affine_param = False,
361
365
  affine_param_batch_decay = 0.99,
362
- affine_param_codebook_decay = 0.9
366
+ affine_param_codebook_decay = 0.9,
367
+ use_cosine_sim = False
363
368
  ):
364
369
  super().__init__()
365
- self.transform_input = identity
370
+ self.transform_input = identity if not use_cosine_sim else l2norm
366
371
 
367
372
  self.decay = decay
368
373
  self.ema_update = ema_update
369
374
  self.manual_ema_update = manual_ema_update
370
375
 
371
- init_fn = uniform_init if not kmeans_init else torch.zeros
372
- embed = init_fn(num_codebooks, codebook_size, dim)
376
+ if kmeans_init:
377
+ embed = torch.zeros(num_codebooks, codebook_size, dim)
378
+ else:
379
+ embed = uniform_init(num_codebooks, codebook_size, dim)
380
+
381
+ if use_cosine_sim:
382
+ embed = l2norm(embed)
373
383
 
374
384
  self.codebook_size = codebook_size
375
385
  self.num_codebooks = num_codebooks
@@ -392,7 +402,7 @@ class EuclideanCodebook(Module):
392
402
  self.kmeans_all_reduce_fn = distributed.all_reduce if use_ddp and sync_kmeans else noop
393
403
  self.all_reduce_fn = distributed.all_reduce if use_ddp else noop
394
404
 
395
- self.register_buffer('initted', torch.Tensor([not kmeans_init]))
405
+ self.register_buffer('initted', tensor(not kmeans_init))
396
406
  self.register_buffer('cluster_size', torch.ones(num_codebooks, codebook_size))
397
407
  self.register_buffer('embed_avg', embed.clone())
398
408
 
@@ -402,6 +412,8 @@ class EuclideanCodebook(Module):
402
412
  else:
403
413
  self.register_buffer('embed', embed)
404
414
 
415
+ self.use_cosine_sim = use_cosine_sim
416
+
405
417
  # affine related params
406
418
 
407
419
  self.affine_param = affine_param
@@ -416,9 +428,9 @@ class EuclideanCodebook(Module):
416
428
  self.register_buffer('batch_mean', None)
417
429
  self.register_buffer('batch_variance', None)
418
430
 
419
- self.register_buffer('codebook_mean_needs_init', torch.Tensor([True]))
431
+ self.register_buffer('codebook_mean_needs_init', tensor(True))
420
432
  self.register_buffer('codebook_mean', torch.empty(num_codebooks, 1, dim))
421
- self.register_buffer('codebook_variance_needs_init', torch.Tensor([True]))
433
+ self.register_buffer('codebook_variance_needs_init', tensor(True))
422
434
  self.register_buffer('codebook_variance', torch.empty(num_codebooks, 1, dim))
423
435
 
424
436
  @torch.jit.ignore
@@ -434,6 +446,7 @@ class EuclideanCodebook(Module):
434
446
  data,
435
447
  self.codebook_size,
436
448
  self.kmeans_iters,
449
+ use_cosine_sim = self.use_cosine_sim,
437
450
  sample_fn = self.sample_fn,
438
451
  all_reduce_fn = self.kmeans_all_reduce_fn
439
452
  )
@@ -443,7 +456,7 @@ class EuclideanCodebook(Module):
443
456
  self.embed_avg.data.copy_(embed_sum)
444
457
  self.cluster_size.data.copy_(cluster_size)
445
458
  self.update_ema()
446
- self.initted.data.copy_(torch.Tensor([True]))
459
+ self.initted.data.copy_(tensor(True))
447
460
 
448
461
  @torch.jit.ignore
449
462
  def update_with_decay(self, buffer_name, new_value, decay):
@@ -452,7 +465,7 @@ class EuclideanCodebook(Module):
452
465
  needs_init = getattr(self, buffer_name + "_needs_init", False)
453
466
 
454
467
  if needs_init:
455
- self.register_buffer(buffer_name + "_needs_init", torch.Tensor([False]))
468
+ self.register_buffer(buffer_name + "_needs_init", tensor(False))
456
469
 
457
470
  if not exists(old_value) or needs_init:
458
471
  self.register_buffer(buffer_name, new_value.detach())
@@ -495,7 +508,7 @@ class EuclideanCodebook(Module):
495
508
 
496
509
  # number of vectors, for denominator
497
510
 
498
- num_vectors = torch.tensor([num_vectors], device = device, dtype = dtype)
511
+ num_vectors = tensor(num_vectors, device = device, dtype = dtype)
499
512
  distributed.all_reduce(num_vectors)
500
513
 
501
514
  # calculate distributed mean
@@ -515,6 +528,9 @@ class EuclideanCodebook(Module):
515
528
  self.update_with_decay('batch_variance', batch_variance, self.affine_param_batch_decay)
516
529
 
517
530
  def replace(self, batch_samples, batch_mask):
531
+ if self.use_cosine_sim:
532
+ batch_samples = l2norm(batch_samples)
533
+
518
534
  for ind, (samples, mask) in enumerate(zip(batch_samples, batch_mask)):
519
535
  sampled = self.replace_sample_fn(rearrange(samples, '... -> 1 ...'), mask.sum().item())
520
536
  sampled = rearrange(sampled, '1 ... -> ...')
@@ -539,247 +555,74 @@ class EuclideanCodebook(Module):
539
555
  cluster_size = laplace_smoothing(self.cluster_size, self.codebook_size, self.eps) * self.cluster_size.sum(dim = -1, keepdim = True)
540
556
 
541
557
  embed_normalized = self.embed_avg / rearrange(cluster_size, '... -> ... 1')
558
+
559
+ if self.use_cosine_sim:
560
+ embed_normalized = l2norm(embed_normalized)
561
+
542
562
  self.embed.data.copy_(embed_normalized)
543
563
 
544
- @autocast('cuda', enabled = False)
545
- def forward(
564
+ def update_ema_part(
546
565
  self,
547
- x,
548
- sample_codebook_temp = None,
566
+ flatten,
567
+ embed_onehot,
549
568
  mask = None,
550
- freeze_codebook = False,
551
- codebook_transform_fn: Callable | None = None,
552
569
  ema_update_weight: Tensor | Callable | None = None,
553
570
  accum_ema_update = False
554
571
  ):
555
- needs_codebook_dim = x.ndim < 4
556
- sample_codebook_temp = default(sample_codebook_temp, self.sample_codebook_temp)
557
-
558
- x = x.float()
559
-
560
- if needs_codebook_dim:
561
- x = rearrange(x, '... -> 1 ...')
562
-
563
- dtype = x.dtype
564
- flatten, unpack_one = pack_one(x, 'h * d')
565
-
566
- if exists(mask):
567
- mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
568
-
569
- self.init_embed_(flatten, mask = mask)
570
-
571
- if self.affine_param:
572
- self.update_affine(flatten, self.embed, mask = mask)
573
-
574
- # get maybe learnable codes
575
-
576
- embed = self.embed if self.learnable_codebook else self.embed.detach()
577
-
578
- embed = embed.to(dtype)
579
-
580
- # affine params
581
572
  if self.affine_param:
582
573
  codebook_std = self.codebook_variance.clamp(min = 1e-5).sqrt()
583
574
  batch_std = self.batch_variance.clamp(min = 1e-5).sqrt()
584
- embed = (embed - self.codebook_mean) * (batch_std / codebook_std) + self.batch_mean
585
-
586
- # handle maybe implicit neural codebook
587
- # and calculate distance
575
+ flatten = (flatten - self.batch_mean) * (codebook_std / batch_std) + self.codebook_mean
588
576
 
589
- if exists(codebook_transform_fn):
590
- transformed_embed = codebook_transform_fn(embed)
591
- transformed_embed = rearrange(transformed_embed, 'h b n c d -> h (b n) c d')
592
- broadcastable_input = rearrange(flatten, '... d -> ... 1 d')
593
-
594
- dist = -F.pairwise_distance(broadcastable_input, transformed_embed)
595
- else:
596
- dist = -cdist(flatten, embed)
597
-
598
- # sample or argmax depending on temperature
599
-
600
- embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, temperature = sample_codebook_temp, training = self.training)
601
-
602
- embed_ind = unpack_one(embed_ind, 'h *')
577
+ if exists(mask):
578
+ embed_onehot[~mask] = 0.
603
579
 
604
- if exists(codebook_transform_fn):
605
- transformed_embed = unpack_one(transformed_embed, 'h * c d')
580
+ cluster_size = embed_onehot.sum(dim = 1)
581
+ self.all_reduce_fn(cluster_size)
606
582
 
607
- if self.training:
608
- unpacked_onehot = unpack_one(embed_onehot, 'h * c')
583
+ embed_sum = einsum('h n d, h n c -> h c d', flatten, embed_onehot)
584
+ embed_sum = embed_sum.contiguous()
585
+ self.all_reduce_fn(embed_sum)
609
586
 
610
- if exists(codebook_transform_fn):
611
- quantize = einsum('h b n c, h b n c d -> h b n d', unpacked_onehot, transformed_embed)
612
- else:
613
- quantize = einsum('h b n c, h c d -> h b n d', unpacked_onehot, embed)
587
+ if callable(ema_update_weight):
588
+ ema_update_weight = ema_update_weight(embed_sum, cluster_size)
614
589
 
590
+ if accum_ema_update:
591
+ accum_grad_(self.cluster_size, cluster_size)
592
+ accum_grad_(self.embed_avg, embed_sum)
615
593
  else:
616
- if exists(codebook_transform_fn):
617
- # quantize = einx.get_at('h b n [c] d, h b n -> h b n d', transformed_embed, embed_ind)
618
-
619
- repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n 1 d', d = transformed_embed.shape[-1])
620
- quantize = transformed_embed.gather(-2, repeated_embed_ind)
621
- quantize = rearrange(quantize, 'h b n 1 d -> h b n d')
622
-
623
- else:
624
- # quantize = einx.get_at('h [c] d, h b n -> h b n d', embed, embed_ind)
625
-
626
- repeated_embed = repeat(embed, 'h c d -> h b c d', b = embed_ind.shape[1])
627
- repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
628
- quantize = repeated_embed.gather(-2, repeated_embed_ind)
629
-
630
- if self.training and self.ema_update and not freeze_codebook:
631
-
632
- if self.affine_param:
633
- flatten = (flatten - self.batch_mean) * (codebook_std / batch_std) + self.codebook_mean
634
-
635
- if exists(mask):
636
- embed_onehot[~mask] = 0.
637
-
638
- cluster_size = embed_onehot.sum(dim = 1)
639
- self.all_reduce_fn(cluster_size)
594
+ ema_inplace(self.cluster_size, cluster_size, self.decay, ema_update_weight)
595
+ ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
640
596
 
641
- embed_sum = einsum('h n d, h n c -> h c d', flatten, embed_onehot)
642
- embed_sum = embed_sum.contiguous()
643
- self.all_reduce_fn(embed_sum)
597
+ if not self.manual_ema_update:
598
+ self.update_ema()
599
+ self.expire_codes_(flatten)
644
600
 
645
- if callable(ema_update_weight):
646
- ema_update_weight = ema_update_weight(embed_sum, cluster_size)
647
-
648
- if accum_ema_update:
649
- accum_grad_(self.cluster_size, cluster_size)
650
- accum_grad_(self.embed_avg, embed_sum)
651
- else:
652
- ema_inplace(self.cluster_size, cluster_size, self.decay, ema_update_weight)
653
- ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
654
-
655
- if not self.manual_ema_update:
656
- self.update_ema()
657
- self.expire_codes_(x)
658
-
659
- if needs_codebook_dim:
660
- quantize, embed_ind = map(lambda t: rearrange(t, '1 ... -> ...'), (quantize, embed_ind))
661
-
662
- dist = unpack_one(dist, 'h * d')
663
-
664
- return quantize, embed_ind, dist
665
-
666
- class CosineSimCodebook(Module):
667
- def __init__(
601
+ def update_ema_indices(
668
602
  self,
669
- dim,
670
- codebook_size,
671
- num_codebooks = 1,
672
- kmeans_init = False,
673
- kmeans_iters = 10,
674
- sync_kmeans = True,
675
- decay = 0.8,
676
- eps = 1e-5,
677
- threshold_ema_dead_code = 2,
678
- reset_cluster_size = None,
679
- use_ddp = False,
680
- learnable_codebook = False,
681
- gumbel_sample = gumbel_sample,
682
- sample_codebook_temp = 1.,
683
- ema_update = True,
684
- manual_ema_update = False
603
+ x,
604
+ embed_ind,
605
+ mask = None,
606
+ ema_update_weight: Tensor | Callable | None = None,
607
+ accum_ema_update = False
685
608
  ):
686
- super().__init__()
687
- self.transform_input = l2norm
688
-
689
- self.ema_update = ema_update
690
- self.manual_ema_update = manual_ema_update
691
-
692
- self.decay = decay
693
-
694
- if not kmeans_init:
695
- embed = l2norm(uniform_init(num_codebooks, codebook_size, dim))
696
- else:
697
- embed = torch.zeros(num_codebooks, codebook_size, dim)
698
-
699
- self.codebook_size = codebook_size
700
- self.num_codebooks = num_codebooks
701
-
702
- self.kmeans_iters = kmeans_iters
703
- self.eps = eps
704
- self.threshold_ema_dead_code = threshold_ema_dead_code
705
- self.reset_cluster_size = default(reset_cluster_size, threshold_ema_dead_code)
706
-
707
- assert callable(gumbel_sample)
708
- self.gumbel_sample = gumbel_sample
709
- self.sample_codebook_temp = sample_codebook_temp
710
-
711
- self.sample_fn = sample_vectors_distributed if use_ddp and sync_kmeans else batched_sample_vectors
712
-
713
- self.replace_sample_fn = sample_vectors_distributed if use_ddp and sync_kmeans else batched_sample_vectors
714
-
715
- self.kmeans_all_reduce_fn = distributed.all_reduce if use_ddp and sync_kmeans else noop
716
- self.all_reduce_fn = distributed.all_reduce if use_ddp else noop
717
-
718
- self.register_buffer('initted', torch.Tensor([not kmeans_init]))
719
- self.register_buffer('cluster_size', torch.ones(num_codebooks, codebook_size))
720
- self.register_buffer('embed_avg', embed.clone())
609
+ needs_codebook_dim = x.ndim < 4
610
+ x = x.float()
721
611
 
722
- self.learnable_codebook = learnable_codebook
723
- if learnable_codebook:
724
- self.embed = nn.Parameter(embed)
725
- else:
726
- self.register_buffer('embed', embed)
612
+ if needs_codebook_dim:
613
+ x = rearrange(x, '... -> 1 ...')
727
614
 
728
- @torch.jit.ignore
729
- def init_embed_(self, data, mask = None):
730
- if self.initted:
731
- return
615
+ dtype = x.dtype
616
+ flatten, unpack_one = pack_one(x, 'h * d')
732
617
 
733
618
  if exists(mask):
734
- c = data.shape[0]
735
- data = rearrange(data[mask], '(c n) d -> c n d', c = c)
736
-
737
- embed, cluster_size = kmeans(
738
- data,
739
- self.codebook_size,
740
- self.kmeans_iters,
741
- use_cosine_sim = True,
742
- sample_fn = self.sample_fn,
743
- all_reduce_fn = self.kmeans_all_reduce_fn
744
- )
745
-
746
- embed_sum = embed * rearrange(cluster_size, '... -> ... 1')
747
-
748
- self.embed_avg.data.copy_(embed_sum)
749
- self.cluster_size.data.copy_(cluster_size)
750
- self.update_ema()
751
- self.initted.data.copy_(torch.Tensor([True]))
752
-
753
- def replace(self, batch_samples, batch_mask):
754
- batch_samples = l2norm(batch_samples)
755
-
756
- for ind, (samples, mask) in enumerate(zip(batch_samples, batch_mask)):
757
- sampled = self.replace_sample_fn(rearrange(samples, '... -> 1 ...'), mask.sum().item())
758
- sampled = rearrange(sampled, '1 ... -> ...')
759
-
760
- self.embed.data[ind][mask] = sampled
761
- self.embed_avg.data[ind][mask] = sampled * self.reset_cluster_size
762
- self.cluster_size.data[ind][mask] = self.reset_cluster_size
763
-
764
- def expire_codes_(self, batch_samples):
765
- if self.threshold_ema_dead_code == 0:
766
- return
767
-
768
- expired_codes = self.cluster_size < self.threshold_ema_dead_code
769
-
770
- if not torch.any(expired_codes):
771
- return
772
-
773
- batch_samples = rearrange(batch_samples, 'h ... d -> h (...) d')
774
- self.replace(batch_samples, batch_mask = expired_codes)
619
+ mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
775
620
 
776
- def update_ema(self):
777
- cluster_size = laplace_smoothing(self.cluster_size, self.codebook_size, self.eps) * self.cluster_size.sum(dim = -1, keepdim = True)
621
+ embed_ind, _ = pack([embed_ind], 'h *')
622
+ embed_ind = embed_ind.masked_fill(embed_ind == -1, 0)
623
+ embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
778
624
 
779
- embed_normalized = self.embed_avg / rearrange(cluster_size, '... -> ... 1')
780
- embed_normalized = l2norm(embed_normalized)
781
-
782
- self.embed.data.copy_(embed_normalized)
625
+ self.update_ema_part(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
783
626
 
784
627
  @autocast('cuda', enabled = False)
785
628
  def forward(
@@ -789,9 +632,13 @@ class CosineSimCodebook(Module):
789
632
  mask = None,
790
633
  freeze_codebook = False,
791
634
  codebook_transform_fn: Callable | None = None,
792
- ema_update_weight: Tensor | None = None,
793
- accum_ema_update = False
635
+ ema_update_weight: Tensor | Callable | None = None,
636
+ accum_ema_update = False,
637
+ ema_update = None,
638
+ topk = None
794
639
  ):
640
+ ema_update = default(ema_update, self.ema_update)
641
+
795
642
  needs_codebook_dim = x.ndim < 4
796
643
  sample_codebook_temp = default(sample_codebook_temp, self.sample_codebook_temp)
797
644
 
@@ -801,7 +648,6 @@ class CosineSimCodebook(Module):
801
648
  x = rearrange(x, '... -> 1 ...')
802
649
 
803
650
  dtype = x.dtype
804
-
805
651
  flatten, unpack_one = pack_one(x, 'h * d')
806
652
 
807
653
  if exists(mask):
@@ -809,35 +655,63 @@ class CosineSimCodebook(Module):
809
655
 
810
656
  self.init_embed_(flatten, mask = mask)
811
657
 
658
+ if self.affine_param:
659
+ self.update_affine(flatten, self.embed, mask = mask)
660
+
661
+ # get maybe learnable codes
662
+
812
663
  embed = self.embed if self.learnable_codebook else self.embed.detach()
813
664
 
814
665
  embed = embed.to(dtype)
815
666
 
667
+ # affine params
668
+
669
+ if self.affine_param:
670
+ codebook_std = self.codebook_variance.clamp(min = 1e-5).sqrt()
671
+ batch_std = self.batch_variance.clamp(min = 1e-5).sqrt()
672
+ embed = (embed - self.codebook_mean) * (batch_std / codebook_std) + self.batch_mean
673
+
816
674
  # handle maybe implicit neural codebook
817
- # and compute cosine sim distance
675
+ # and calculate distance
818
676
 
819
677
  if exists(codebook_transform_fn):
820
678
  transformed_embed = codebook_transform_fn(embed)
821
679
  transformed_embed = rearrange(transformed_embed, 'h b n c d -> h (b n) c d')
822
- transformed_embed = l2norm(transformed_embed)
823
680
 
824
- dist = einsum('h n d, h n c d -> h n c', flatten, transformed_embed)
681
+ if self.use_cosine_sim:
682
+ transformed_embed = l2norm(transformed_embed)
683
+ dist = einsum('h n d, h n c d -> h n c', flatten, transformed_embed)
684
+ else:
685
+ broadcastable_input = rearrange(flatten, '... d -> ... 1 d')
686
+ dist = -F.pairwise_distance(broadcastable_input, transformed_embed)
825
687
  else:
826
- dist = einsum('h n d, h c d -> h n c', flatten, embed)
688
+ if self.use_cosine_sim:
689
+ dist = einsum('h n d, h c d -> h n c', flatten, embed)
690
+ else:
691
+ dist = -cdist(flatten, embed)
692
+
693
+ # sample or argmax depending on temperature
827
694
 
828
- embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, temperature = sample_codebook_temp, training = self.training)
829
- embed_ind = unpack_one(embed_ind, 'h *')
695
+ embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, topk = topk, temperature = sample_codebook_temp, training = self.training)
696
+
697
+ if exists(topk):
698
+ embed_ind = unpack_one(embed_ind, 'h * k')
699
+ else:
700
+ embed_ind = unpack_one(embed_ind, 'h *')
830
701
 
831
702
  if exists(codebook_transform_fn):
832
703
  transformed_embed = unpack_one(transformed_embed, 'h * c d')
833
704
 
834
705
  if self.training:
835
- unpacked_onehot = unpack_one(embed_onehot, 'h * c')
706
+ if exists(topk):
707
+ unpacked_onehot = unpack_one(embed_onehot, 'h * k c')
708
+ else:
709
+ unpacked_onehot = unpack_one(embed_onehot, 'h * c')
836
710
 
837
711
  if exists(codebook_transform_fn):
838
- quantize = einsum('h b n c, h b n c d -> h b n d', unpacked_onehot, transformed_embed)
712
+ quantize = einsum('h b n ... c, h b n c d -> h b n ... d', unpacked_onehot, transformed_embed)
839
713
  else:
840
- quantize = einsum('h b n c, h c d -> h b n d', unpacked_onehot, embed)
714
+ quantize = einsum('h b n ... c, h c d -> h b n ... d', unpacked_onehot, embed)
841
715
 
842
716
  else:
843
717
  if exists(codebook_transform_fn):
@@ -854,36 +728,14 @@ class CosineSimCodebook(Module):
854
728
  repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
855
729
  quantize = repeated_embed.gather(-2, repeated_embed_ind)
856
730
 
857
- if self.training and self.ema_update and not freeze_codebook:
858
- if exists(mask):
859
- embed_onehot[~mask] = 0.
860
-
861
- bins = embed_onehot.sum(dim = 1)
862
- self.all_reduce_fn(bins)
863
-
864
- embed_sum = einsum('h n d, h n c -> h c d', flatten, embed_onehot)
865
- embed_sum = embed_sum.contiguous()
866
- self.all_reduce_fn(embed_sum)
867
-
868
- if callable(ema_update_weight):
869
- ema_update_weight = ema_update_weight(embed_sum, bins)
870
-
871
- if accum_ema_update:
872
- accum_grad_(self.cluster_size, bins)
873
- accum_grad_(self.embed_avg, embed_sum)
874
- else:
875
-
876
- ema_inplace(self.cluster_size, bins, self.decay, ema_update_weight)
877
- ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
878
-
879
- if not self.manual_ema_update:
880
- self.update_ema()
881
- self.expire_codes_(x)
731
+ if self.training and ema_update and not freeze_codebook and not exists(topk):
732
+ self.update_ema_part(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
882
733
 
883
734
  if needs_codebook_dim:
884
735
  quantize, embed_ind = map(lambda t: rearrange(t, '1 ... -> ...'), (quantize, embed_ind))
885
736
 
886
737
  dist = unpack_one(dist, 'h * d')
738
+
887
739
  return quantize, embed_ind, dist
888
740
 
889
741
  # main class
@@ -1001,7 +853,7 @@ class VectorQuantize(Module):
1001
853
 
1002
854
  self.sync_update_v = sync_update_v
1003
855
 
1004
- codebook_class = EuclideanCodebook if not use_cosine_sim else CosineSimCodebook
856
+ codebook_class = Codebook
1005
857
 
1006
858
  gumbel_sample_fn = partial(
1007
859
  gumbel_sample,
@@ -1027,7 +879,8 @@ class VectorQuantize(Module):
1027
879
  sample_codebook_temp = sample_codebook_temp,
1028
880
  gumbel_sample = gumbel_sample_fn,
1029
881
  ema_update = ema_update,
1030
- manual_ema_update = manual_ema_update
882
+ manual_ema_update = manual_ema_update,
883
+ use_cosine_sim = use_cosine_sim
1031
884
  )
1032
885
 
1033
886
  if affine_param:
@@ -1051,7 +904,7 @@ class VectorQuantize(Module):
1051
904
  self.accept_image_fmap = accept_image_fmap
1052
905
  self.channel_last = channel_last
1053
906
 
1054
- self.register_buffer('zero', torch.tensor(0.), persistent = False)
907
+ self.register_buffer('zero', tensor(0.), persistent = False)
1055
908
 
1056
909
  # for variable lengthed sequences, whether to take care of masking out the padding to 0 (or return the original input)
1057
910
  self.return_zeros_for_masked_padding = return_zeros_for_masked_padding
@@ -1059,6 +912,10 @@ class VectorQuantize(Module):
1059
912
  # whether to freeze the codebook, can be overridden on forward
1060
913
  self.freeze_codebook = freeze_codebook
1061
914
 
915
+ @property
916
+ def ema_update(self):
917
+ return self._codebook.ema_update
918
+
1062
919
  @property
1063
920
  def codebook(self):
1064
921
  codebook = self._codebook.embed
@@ -1133,18 +990,47 @@ class VectorQuantize(Module):
1133
990
  x = self.maybe_split_heads_from_input(x)
1134
991
  self._codebook.expire_codes_(x)
1135
992
 
993
+ def update_ema_indices(self, x, indices, mask = None):
994
+ if self.accept_image_fmap:
995
+ assert not exists(mask)
996
+ height, width = x.shape[-2:]
997
+ x = rearrange(x, 'b c h w -> b (h w) c')
998
+
999
+ if not self.channel_last and not self.accept_image_fmap:
1000
+ x = rearrange(x, 'b d n -> b n d')
1001
+
1002
+ x = self.project_in(x)
1003
+ x = self.maybe_split_heads_from_input(x)
1004
+ x = self._codebook.transform_input(x)
1005
+
1006
+ if self.heads > 1:
1007
+ if self.separate_codebook_per_head:
1008
+ indices = rearrange(indices, 'b n h -> h b n')
1009
+ else:
1010
+ indices = rearrange(indices, 'b n h -> 1 (b h) n')
1011
+
1012
+ if self.accept_image_fmap:
1013
+ indices = rearrange(indices, 'b h w ... -> b (h w) ...')
1014
+
1015
+ if x.ndim == 2: # only one token
1016
+ indices = rearrange(indices, 'b ... -> b 1 ...')
1017
+
1018
+ self._codebook.update_ema_indices(x, indices, mask = mask)
1019
+
1136
1020
  def forward(
1137
1021
  self,
1138
1022
  x,
1139
1023
  indices = None,
1140
1024
  mask = None,
1141
1025
  lens = None,
1026
+ topk = None,
1142
1027
  sample_codebook_temp = None,
1143
1028
  freeze_codebook = None,
1144
1029
  return_loss_breakdown = False,
1145
1030
  codebook_transform_fn: Callable | None = None,
1146
1031
  ema_update_weight: Tensor | None = None,
1147
- accum_ema_update = False
1032
+ accum_ema_update = False,
1033
+ ema_update = None
1148
1034
  ):
1149
1035
  orig_input, input_requires_grad = x, x.requires_grad
1150
1036
 
@@ -1202,7 +1088,9 @@ class VectorQuantize(Module):
1202
1088
  freeze_codebook = freeze_codebook,
1203
1089
  codebook_transform_fn = codebook_transform_fn,
1204
1090
  ema_update_weight = ema_update_weight,
1205
- accum_ema_update = accum_ema_update
1091
+ accum_ema_update = accum_ema_update,
1092
+ ema_update = ema_update and not exists(topk),
1093
+ topk = topk
1206
1094
  )
1207
1095
 
1208
1096
  # quantize
@@ -1301,7 +1189,7 @@ class VectorQuantize(Module):
1301
1189
 
1302
1190
  # aggregate loss
1303
1191
 
1304
- loss = torch.tensor([0.], device = device, requires_grad = self.training)
1192
+ loss = tensor(0., device = device, requires_grad = self.training)
1305
1193
 
1306
1194
  if self.training:
1307
1195
  # calculate codebook diversity loss (negative of entropy) if needed
@@ -1326,9 +1214,19 @@ class VectorQuantize(Module):
1326
1214
 
1327
1215
  commit_loss = calculate_ce_loss(embed_ind)
1328
1216
  else:
1329
- if exists(mask):
1217
+ if exists(topk):
1218
+ # handle special case when returning topk
1219
+
1220
+ repeated_input = repeat(orig_input, '... d -> ... k d', k = topk)
1221
+ commit_loss = F.mse_loss(commit_quantize, repeated_input, reduction = 'none')
1222
+ commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
1223
+
1224
+ if exists(mask):
1225
+ commit_loss = einx.where('..., ... k, -> ... k', mask, commit_loss, 0.)
1226
+
1227
+ elif exists(mask):
1330
1228
  # with variable lengthed sequences
1331
- commit_loss = F.mse_loss(commit_quantize, x, reduction = 'none')
1229
+ commit_loss = F.mse_loss(commit_quantize, orig_input, reduction = 'none')
1332
1230
 
1333
1231
  loss_mask = mask
1334
1232
  if is_multiheaded:
@@ -1391,7 +1289,7 @@ class VectorQuantize(Module):
1391
1289
  masked_out_value = torch.zeros_like(orig_input)
1392
1290
 
1393
1291
  quantize = einx.where(
1394
- 'b n, b n d, b n d -> b n d',
1292
+ 'b n, b n ... d, b n d -> b n ... d',
1395
1293
  mask,
1396
1294
  quantize,
1397
1295
  masked_out_value