vector-quantize-pytorch 1.31.0__tar.gz → 1.31.2__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 (21) hide show
  1. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/PKG-INFO +3 -2
  2. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/pyproject.toml +3 -2
  3. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/__init__.py +1 -0
  4. vector_quantize_pytorch-1.31.2/vector_quantize_pytorch/evo_lookup_free_quantization.py +289 -0
  5. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +57 -30
  6. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/.gitignore +0 -0
  7. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/LICENSE +0 -0
  8. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/README.md +0 -0
  9. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/binary_mapper.py +0 -0
  10. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/finite_scalar_perturbation.py +0 -0
  11. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  12. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/hierarchical_vq.py +0 -0
  13. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
  14. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  15. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  16. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
  17. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
  18. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  19. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_vq.py +0 -0
  20. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/sim_vq.py +0 -0
  21. {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/utils.py +0 -0
@@ -1,6 +1,6 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.31.0
3
+ Version: 1.31.2
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,6 +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-einops-utils>=0.1.12
39
40
  Requires-Dist: torch>=2.4
40
41
  Provides-Extra: examples
41
42
  Requires-Dist: torchvision; extra == 'examples'
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.31.0"
3
+ version = "1.31.2"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -23,9 +23,10 @@ classifiers=[
23
23
  ]
24
24
 
25
25
  dependencies = [
26
- "torch>=2.4",
27
26
  "einops>=0.8.0",
28
27
  "einx>=0.3.0",
28
+ "torch>=2.4",
29
+ "torch-einops-utils>=0.1.12",
29
30
  ]
30
31
 
31
32
  [project.urls]
@@ -4,6 +4,7 @@ from vector_quantize_pytorch.random_projection_quantizer import RandomProjection
4
4
  from vector_quantize_pytorch.finite_scalar_quantization import FSQ
5
5
  from vector_quantize_pytorch.finite_scalar_perturbation import FSP
6
6
  from vector_quantize_pytorch.lookup_free_quantization import LFQ
7
+ from vector_quantize_pytorch.evo_lookup_free_quantization import EvoLFQ
7
8
  from vector_quantize_pytorch.residual_lfq import ResidualLFQ, GroupedResidualLFQ
8
9
  from vector_quantize_pytorch.residual_fsq import ResidualFSQ, GroupedResidualFSQ
9
10
  from vector_quantize_pytorch.latent_quantization import LatentQuantize
@@ -0,0 +1,289 @@
1
+ from __future__ import annotations
2
+ from collections import namedtuple
3
+
4
+ import torch
5
+ from torch import nn, Tensor, is_tensor
6
+ from torch.nn import Module
7
+ import torch.nn.functional as F
8
+
9
+ from einops import rearrange, repeat, reduce
10
+ from torch_einops_utils import pack_with_inverse, temp_eval
11
+
12
+ from vector_quantize_pytorch.lookup_free_quantization import LFQ
13
+
14
+ # constants
15
+
16
+ Return = namedtuple('Return', ['reconstructed', 'indices', 'entropy_aux_loss'])
17
+ Result = namedtuple('Result', ['pop_bits', 'best_gene', 'best_fitness', 'best_decoded'])
18
+
19
+ # helper functions
20
+
21
+ def exists(val):
22
+ return val is not None
23
+
24
+ def default(val, d):
25
+ return val if exists(val) else d
26
+
27
+ # main class
28
+
29
+ class EvoLFQ(Module):
30
+ def __init__(
31
+ self,
32
+ encoder: Module,
33
+ decoder: Module,
34
+ lfq: LFQ | None = None,
35
+ *,
36
+ dim: int | None = None,
37
+ codebook_size: int | None = None,
38
+ num_codebooks: int = 1,
39
+ pop_size: int = 64,
40
+ mutation_rate: float = 0.02,
41
+ tournament_size: int = 2,
42
+ elitism_count: int = 1,
43
+ generations: int = 50,
44
+ **lfq_kwargs
45
+ ):
46
+ super().__init__()
47
+ assert pop_size > elitism_count, 'pop_size must be greater than elitism_count'
48
+
49
+ self.encoder = encoder
50
+ self.decoder = decoder
51
+
52
+ if not exists(lfq):
53
+ assert exists(dim) or exists(codebook_size), 'either lfq instance or dim / codebook_size must be supplied to EvoLFQ'
54
+ lfq = LFQ(
55
+ dim = dim,
56
+ codebook_size = codebook_size,
57
+ num_codebooks = num_codebooks,
58
+ **lfq_kwargs
59
+ )
60
+
61
+ self.lfq = lfq
62
+ self.pop_size = pop_size
63
+ self.mutation_rate = mutation_rate
64
+ self.tournament_size = tournament_size
65
+ self.elitism_count = elitism_count
66
+ self.generations = generations
67
+
68
+ self.register_buffer('zero', torch.tensor(0.), persistent = False)
69
+
70
+ @property
71
+ def device(self):
72
+ return self.zero.device
73
+
74
+ def forward(self, x, **kwargs):
75
+ latents = self.encoder(x)
76
+ is_2d = latents.ndim == 2
77
+
78
+ if is_2d:
79
+ latents = rearrange(latents, 'b d -> b 1 d')
80
+
81
+ quantized, indices, aux_loss = self.lfq(latents, **kwargs)
82
+
83
+ if is_2d:
84
+ quantized = rearrange(quantized, 'b 1 d -> b d')
85
+ indices = rearrange(indices, 'b 1 ... -> b ...')
86
+
87
+ reconstructed = self.decoder(quantized)
88
+ return Return(reconstructed, indices, aux_loss)
89
+
90
+ @torch.no_grad()
91
+ def encode(self, x, return_signs = False):
92
+ with temp_eval(self):
93
+ latents = self.encoder(x)
94
+ is_2d = latents.ndim == 2
95
+
96
+ if is_2d:
97
+ latents = rearrange(latents, 'b d -> b 1 d')
98
+
99
+ quantized, indices, _ = self.lfq(latents)
100
+
101
+ if is_2d:
102
+ quantized = rearrange(quantized, 'b 1 d -> b d')
103
+
104
+ if return_signs:
105
+ return torch.where(quantized > 0, 1.0, -1.0)
106
+
107
+ return (quantized > 0).float()
108
+
109
+ @torch.no_grad()
110
+ def decode_bits(self, bits):
111
+ """
112
+ Converts binary bits (0/1 or -1/+1) to codebook indices and uses lfq.indices_to_codes
113
+ to decode back to data space.
114
+ """
115
+ with temp_eval(self):
116
+ bits = (bits > 0).int()
117
+ bits, inverse = pack_with_inverse(bits, '* d')
118
+
119
+ codebook_dim = self.lfq.codebook_dim
120
+ num_codebooks = self.lfq.num_codebooks
121
+
122
+ if num_codebooks > 1 and bits.shape[-1] == codebook_dim * num_codebooks:
123
+ bits = rearrange(bits, 'b (c d) -> b c d', c = num_codebooks, d = codebook_dim)
124
+
125
+ if bits.ndim == 2:
126
+ bits = rearrange(bits, 'b d -> b 1 d')
127
+
128
+ mask = 2 ** torch.arange(codebook_dim - 1, -1, -1, device = self.device)
129
+ indices = reduce(bits * mask, '... d -> ...', 'sum')
130
+
131
+ if not self.lfq.keep_num_codebooks_dim and indices.ndim >= 2 and indices.shape[-1] == 1:
132
+ indices = rearrange(indices, '... 1 -> ...')
133
+
134
+ codes = self.lfq.indices_to_codes(indices)
135
+
136
+ if codes.ndim == 3 and codes.shape[1] == 1:
137
+ codes = rearrange(codes, 'b 1 d -> b d')
138
+
139
+ codes = inverse(codes, '* d')
140
+ return self.decoder(codes)
141
+
142
+ # genetic algorithm helpers
143
+
144
+ def init_random_population(self, pop_size = None, shape = None, device = None, is_sign = False):
145
+ pop_size = default(pop_size, self.pop_size)
146
+ device = default(device, self.device)
147
+ bits = torch.randint(0, 2, (pop_size, *shape), device = device).float()
148
+
149
+ if is_sign:
150
+ return torch.where(bits > 0, 1.0, -1.0)
151
+
152
+ return bits
153
+
154
+ def init_population_from_data(self, x, pop_size = None, mutation_rate = None, is_sign = False):
155
+ pop_size = default(pop_size, self.pop_size)
156
+ mutation_rate = default(mutation_rate, self.mutation_rate)
157
+ bits = self.encode(x, return_signs = is_sign)
158
+
159
+ num_samples, *_ = bits.shape
160
+
161
+ if num_samples < pop_size:
162
+ repeats = (pop_size + num_samples - 1) // num_samples
163
+ bits = repeat(bits, 'b ... -> (r b) ...', r = repeats)[:pop_size]
164
+ else:
165
+ bits = bits[:pop_size]
166
+
167
+ if mutation_rate > 0:
168
+ bits = self.mutate(bits, mutation_rate = mutation_rate, is_sign = is_sign)
169
+
170
+ return bits
171
+
172
+ def uniform_crossover(self, parent1, parent2):
173
+ crossover_mask = torch.rand_like(parent1.float()) < 0.5
174
+ return torch.where(crossover_mask, parent1, parent2)
175
+
176
+ def mutate(self, population, mutation_rate = None, is_sign = False):
177
+ mutation_rate = default(mutation_rate, self.mutation_rate)
178
+ flip_mask = torch.rand_like(population.float()) < mutation_rate
179
+
180
+ if is_sign:
181
+ return torch.where(flip_mask, -population, population)
182
+
183
+ return torch.where(flip_mask, 1.0 - population, population)
184
+
185
+ def tournament_selection(self, population, fitnesses, tournament_size = None):
186
+ tournament_size = default(tournament_size, self.tournament_size)
187
+ pop_size, *_ = population.shape
188
+
189
+ contenders = torch.randint(0, pop_size, (pop_size, tournament_size), device = self.device)
190
+ winner_indices = contenders.gather(1, fitnesses[contenders].argmax(dim = -1, keepdim = True))
191
+ return population[rearrange(winner_indices, 'b 1 -> b')]
192
+
193
+ @torch.no_grad()
194
+ def step(
195
+ self,
196
+ pop_bits,
197
+ fitness_fn,
198
+ tournament_size = None,
199
+ mutation_rate = None,
200
+ elitism_count = None,
201
+ is_sign = False,
202
+ batch_size = None,
203
+ **kwargs
204
+ ):
205
+ tournament_size = default(tournament_size, self.tournament_size)
206
+ mutation_rate = default(mutation_rate, self.mutation_rate)
207
+ elitism_count = default(elitism_count, self.elitism_count)
208
+
209
+ pop_size, *_ = pop_bits.shape
210
+ num_offspring = pop_size - elitism_count
211
+
212
+ # evaluate fitness
213
+
214
+ if exists(batch_size):
215
+ decoded_chunks = [self.decode_bits(chunk) for chunk in pop_bits.split(batch_size)]
216
+ decoded = torch.cat(decoded_chunks, dim = 0)
217
+ else:
218
+ decoded = self.decode_bits(pop_bits)
219
+
220
+ fitnesses = fitness_fn(decoded, pop_bits)
221
+ if not is_tensor(fitnesses):
222
+ fitnesses = torch.tensor(fitnesses, device = self.device, dtype = torch.float32)
223
+
224
+ # sort by fitness descending and preserve elites
225
+
226
+ sorted_indices = torch.argsort(fitnesses, descending = True)
227
+ elites = pop_bits[sorted_indices[:elitism_count]].clone()
228
+
229
+ # tournament selection: 2 parent tournaments for each of the num_offspring children needed
230
+
231
+ contenders = torch.randint(0, pop_size, (num_offspring, 2, tournament_size), device = self.device)
232
+ winner_rel_idx = fitnesses[contenders].argmax(dim = -1, keepdim = True)
233
+ parent_indices = rearrange(contenders.gather(-1, winner_rel_idx), 'n p 1 -> n p')
234
+
235
+ # crossover: 1 child produced per 2 parents
236
+
237
+ p1, p2 = rearrange(pop_bits[parent_indices], 'n p ... -> p n ...')
238
+ offspring = self.uniform_crossover(p1, p2)
239
+
240
+ # mutation all at once
241
+
242
+ offspring = self.mutate(offspring, mutation_rate = mutation_rate, is_sign = is_sign)
243
+
244
+ next_pop = torch.cat([elites, offspring], dim = 0)
245
+ return next_pop, fitnesses
246
+
247
+ @torch.no_grad()
248
+ def evolve(
249
+ self,
250
+ fitness_fn,
251
+ pop_bits = None,
252
+ pop_size = None,
253
+ shape = None,
254
+ generations = None,
255
+ is_sign = False,
256
+ return_best_decoded = False,
257
+ **step_kwargs
258
+ ):
259
+ pop_size = default(pop_size, self.pop_size)
260
+ generations = default(generations, self.generations)
261
+
262
+ if not exists(pop_bits):
263
+ assert exists(shape), 'shape must be provided if pop_bits is not supplied to evolve()'
264
+ pop_bits = self.init_random_population(pop_size, shape, is_sign = is_sign)
265
+
266
+ best_fitness = float('-inf')
267
+ best_gene = None
268
+
269
+ for _ in range(generations):
270
+ pop_bits, fitnesses = self.step(
271
+ pop_bits,
272
+ fitness_fn,
273
+ is_sign = is_sign,
274
+ **step_kwargs
275
+ )
276
+
277
+ max_fit_idx = torch.argmax(fitnesses)
278
+ max_fit = fitnesses[max_fit_idx].item()
279
+
280
+ if max_fit > best_fitness:
281
+ best_fitness = max_fit
282
+ best_gene = pop_bits[max_fit_idx].clone()
283
+
284
+ best_decoded = None
285
+ if return_best_decoded:
286
+ best_decoded = self.decode_bits(rearrange(best_gene, '... -> 1 ...'))
287
+ best_decoded = rearrange(best_decoded, '1 ... -> ...')
288
+
289
+ yield Result(pop_bits, best_gene, best_fitness, best_decoded)
@@ -654,14 +654,15 @@ class Codebook(Module):
654
654
 
655
655
  if needs_codebook_dim:
656
656
  x = rearrange(x, '... -> 1 ...')
657
+ embed_ind = rearrange(embed_ind, '... -> 1 ...')
657
658
 
658
659
  dtype = x.dtype
659
- flatten, unpack_one = pack_one(x, 'h * d')
660
+ flatten, _ = pack_one(x, 'h * d')
660
661
 
661
662
  if exists(mask):
662
663
  mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
663
664
 
664
- embed_ind, _ = pack([embed_ind], 'h *')
665
+ embed_ind, _ = pack_one(embed_ind, 'h *')
665
666
  embed_ind = embed_ind.masked_fill(embed_ind == -1, 0)
666
667
  embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
667
668
 
@@ -697,6 +698,8 @@ class Codebook(Module):
697
698
  dtype = x.dtype
698
699
  flatten, unpack_one = pack_one(x, 'h * d')
699
700
 
701
+ is_topk = exists(topk)
702
+
700
703
  if exists(mask):
701
704
  mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
702
705
 
@@ -746,19 +749,13 @@ class Codebook(Module):
746
749
 
747
750
  embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, topk = topk, temperature = sample_codebook_temp, training = self.training)
748
751
 
749
- if exists(topk):
750
- embed_ind = unpack_one(embed_ind, 'h * k')
751
- else:
752
- embed_ind = unpack_one(embed_ind, 'h *')
752
+ embed_ind = unpack_one(embed_ind, 'h * k' if is_topk else 'h *')
753
753
 
754
754
  if exists(codebook_transform_fn):
755
755
  transformed_embed = unpack_one(transformed_embed, 'h * c d')
756
756
 
757
- if self.training:
758
- if exists(topk):
759
- unpacked_onehot = unpack_one(embed_onehot, 'h * k c')
760
- else:
761
- unpacked_onehot = unpack_one(embed_onehot, 'h * c')
757
+ if self.training or is_topk:
758
+ unpacked_onehot = unpack_one(embed_onehot, 'h * k c' if is_topk else 'h * c')
762
759
 
763
760
  if exists(codebook_transform_fn):
764
761
  quantize = einsum('h b n ... c, h b n c d -> h b n ... d', unpacked_onehot, transformed_embed)
@@ -780,7 +777,7 @@ class Codebook(Module):
780
777
  repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
781
778
  quantize = repeated_embed.gather(-2, repeated_embed_ind)
782
779
 
783
- if self.training and update_usage and not freeze_codebook and not exists(topk):
780
+ if self.training and update_usage and not freeze_codebook and not is_topk:
784
781
  self.update_codebook(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update, ema_update = ema_update)
785
782
 
786
783
  if needs_codebook_dim:
@@ -1107,6 +1104,18 @@ class VectorQuantize(Module):
1107
1104
  ):
1108
1105
  orig_input, input_requires_grad = x, x.requires_grad
1109
1106
 
1107
+ is_topk = exists(topk)
1108
+
1109
+ if is_topk:
1110
+ topk = min(int(topk), self.codebook_size)
1111
+
1112
+ assert (
1113
+ self.heads == 1
1114
+ and topk > 0
1115
+ and not (self.commitment_use_cross_entropy_loss and self.training)
1116
+ and not (exists(self.in_place_codebook_optimizer) and self.training)
1117
+ ), 'topk codes only supported for single-head codebooks without cross entropy commitment loss or in-place codebook optimizer'
1118
+
1110
1119
  # freezing codebook
1111
1120
 
1112
1121
  freeze_codebook = default(freeze_codebook, self.freeze_codebook)
@@ -1128,7 +1137,7 @@ class VectorQuantize(Module):
1128
1137
 
1129
1138
  shape, dtype, device, heads, is_multiheaded, codebook_size, return_loss = x.shape, x.dtype, x.device, self.heads, self.heads > 1, self.codebook_size, exists(indices)
1130
1139
 
1131
- need_transpose = not self.channel_last and not self.accept_image_fmap and not self.accept_3d_fmap
1140
+ need_transpose = not self.channel_last and not self.accept_image_fmap and not self.accept_3d_fmap and not only_one
1132
1141
  should_inplace_optimize = exists(self.in_place_codebook_optimizer)
1133
1142
 
1134
1143
  # rearrange inputs
@@ -1167,7 +1176,7 @@ class VectorQuantize(Module):
1167
1176
  codebook_transform_fn = codebook_transform_fn,
1168
1177
  ema_update_weight = ema_update_weight,
1169
1178
  accum_ema_update = accum_ema_update,
1170
- ema_update = ema_update and not exists(topk),
1179
+ ema_update = ema_update and not is_topk,
1171
1180
  topk = topk
1172
1181
  )
1173
1182
 
@@ -1209,17 +1218,14 @@ class VectorQuantize(Module):
1209
1218
  codebook_forward_kwargs.update(update_usage = False)
1210
1219
  quantize, embed_ind, distances = self._codebook(x, **codebook_forward_kwargs)
1211
1220
 
1221
+ if is_topk:
1222
+ x = repeat(x, '... d -> ... k d', k = topk)
1223
+
1212
1224
  if self.training:
1213
- # determine code to use for commitment loss
1214
1225
  maybe_detach = torch.detach if not self.learnable_codebook or freeze_codebook else identity
1215
1226
 
1216
1227
  commit_quantize = maybe_detach(quantize)
1217
1228
 
1218
- # maybe expand input if returning topk codes
1219
-
1220
- if exists(topk):
1221
- x = repeat(x, '... d -> ... k d', k = topk)
1222
-
1223
1229
  # spare rotation trick calculation if inputs do not need gradients
1224
1230
 
1225
1231
  if input_requires_grad and self.route_gradients_to_input:
@@ -1279,7 +1285,19 @@ class VectorQuantize(Module):
1279
1285
 
1280
1286
  # aggregate loss
1281
1287
 
1282
- loss = tensor(0., device = device, requires_grad = self.training)
1288
+ loss = self.zero
1289
+
1290
+ if self.training:
1291
+ loss.requires_grad_()
1292
+
1293
+ if is_topk and not self.training:
1294
+ commit_loss = F.mse_loss(quantize, x, reduction = 'none')
1295
+ commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
1296
+
1297
+ if exists(mask):
1298
+ commit_loss = einx.where('..., ... k, -> ... k', mask, commit_loss, 0.)
1299
+
1300
+ loss = commit_loss * self.commitment_weight if self.has_commitment_loss else commit_loss
1283
1301
 
1284
1302
  if self.training:
1285
1303
  # calculate codebook diversity loss (negative of entropy) if needed
@@ -1304,11 +1322,8 @@ class VectorQuantize(Module):
1304
1322
 
1305
1323
  commit_loss = calculate_ce_loss(embed_ind)
1306
1324
  else:
1307
- if exists(topk):
1308
- # handle special case when returning topk
1309
-
1310
- repeated_input = repeat(orig_input, '... d -> ... k d', k = topk)
1311
- commit_loss = F.mse_loss(commit_quantize, repeated_input, reduction = 'none')
1325
+ if is_topk:
1326
+ commit_loss = F.mse_loss(commit_quantize, x, reduction = 'none')
1312
1327
  commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
1313
1328
 
1314
1329
  if exists(mask):
@@ -1347,6 +1362,18 @@ class VectorQuantize(Module):
1347
1362
  orthogonal_reg_loss = orthogonal_loss_fn(codebook)
1348
1363
  loss = loss + orthogonal_reg_loss * self.orthogonal_reg_weight
1349
1364
 
1365
+ # reshape the topk loss to have the same shape as the embedding indices, if it is a per token loss
1366
+
1367
+ if is_topk and loss.ndim > 1:
1368
+ if only_one:
1369
+ loss = rearrange(loss, 'b 1 ... -> b ...')
1370
+
1371
+ if self.accept_image_fmap:
1372
+ loss = rearrange(loss, 'b (h w) ... -> b h w ...', h = height, w = width)
1373
+
1374
+ if self.accept_3d_fmap:
1375
+ loss = rearrange(loss, "b (d h w) ... -> b d h w ...", d = depth, h = height, w = width)
1376
+
1350
1377
  # handle multi-headed quantized embeddings
1351
1378
 
1352
1379
  if is_multiheaded:
@@ -1362,16 +1389,16 @@ class VectorQuantize(Module):
1362
1389
  # rearrange quantized embeddings
1363
1390
 
1364
1391
  if need_transpose:
1365
- quantize = rearrange(quantize, 'b n d -> b d n')
1392
+ quantize = rearrange(quantize, 'b n ... d -> b d n ...')
1366
1393
 
1367
1394
  if self.accept_image_fmap:
1368
- quantize = rearrange(quantize, 'b (h w) c -> b c h w', h = height, w = width)
1395
+ quantize = rearrange(quantize, 'b (h w) ... c -> b c h w ...', h = height, w = width)
1369
1396
 
1370
1397
  if self.accept_3d_fmap:
1371
- quantize = rearrange(quantize, "b (d h w) c -> b c d h w", d=depth, h=height, w=width)
1398
+ quantize = rearrange(quantize, "b (d h w) ... c -> b c d h w ...", d=depth, h=height, w=width)
1372
1399
 
1373
1400
  if only_one:
1374
- quantize = rearrange(quantize, 'b 1 d -> b d')
1401
+ quantize = rearrange(quantize, 'b 1 ... d -> b ... d')
1375
1402
 
1376
1403
  # if masking, only return quantized for where mask has True
1377
1404