vector-quantize-pytorch 1.31.2__tar.gz → 1.31.4__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.2 → vector_quantize_pytorch-1.31.4}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_vq.py +22 -1
  4. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/vector_quantize_pytorch.py +21 -17
  5. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/.gitignore +0 -0
  6. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/LICENSE +0 -0
  7. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/README.md +0 -0
  8. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/__init__.py +0 -0
  9. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/binary_mapper.py +0 -0
  10. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/evo_lookup_free_quantization.py +0 -0
  11. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/finite_scalar_perturbation.py +0 -0
  12. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  13. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/hierarchical_vq.py +0 -0
  14. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
  15. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  16. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  17. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
  18. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
  19. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  20. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/sim_vq.py +0 -0
  21. {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.31.2
3
+ Version: 1.31.4
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.31.2"
3
+ version = "1.31.4"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -185,6 +185,7 @@ class ResidualVQ(Module):
185
185
  eval_beam_size = None,
186
186
  beam_score_quantizer_weights: list[float] | None = None,
187
187
  quant_grad_frac = 0.,
188
+ return_zeros_for_masked_padding = True,
188
189
  **vq_kwargs
189
190
  ):
190
191
  super().__init__()
@@ -202,6 +203,9 @@ class ResidualVQ(Module):
202
203
 
203
204
  self.accept_image_fmap = accept_image_fmap
204
205
 
206
+ self.return_zeros_for_masked_padding = return_zeros_for_masked_padding
207
+ vq_kwargs.update(return_zeros_for_masked_padding = return_zeros_for_masked_padding)
208
+
205
209
  self.implicit_neural_codebook = implicit_neural_codebook
206
210
 
207
211
  if implicit_neural_codebook:
@@ -397,6 +401,8 @@ class ResidualVQ(Module):
397
401
 
398
402
  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
399
403
 
404
+ orig_input = x
405
+
400
406
  beam_size = default(beam_size, self.beam_size if self.training else self.eval_beam_size)
401
407
 
402
408
  is_beam_search = exists(beam_size) and beam_size > 1
@@ -577,7 +583,7 @@ class ResidualVQ(Module):
577
583
 
578
584
  if exists(mask):
579
585
  all_losses = einx.where('..., ... l,', mask, all_losses, 0.)
580
- all_losses = reduce(all_losses, '... l -> l', 'sum') / mask.sum(dim = -1).clamp_min(1e-4)
586
+ all_losses = reduce(all_losses, '... l -> l', 'sum') / mask.sum().clamp_min(1e-4)
581
587
  else:
582
588
  all_losses = reduce(all_losses, '... l -> l', 'mean')
583
589
 
@@ -609,6 +615,21 @@ class ResidualVQ(Module):
609
615
 
610
616
  quantized_out = self.project_out(quantized_out)
611
617
 
618
+ # if masking, only return quantized for where mask has True
619
+
620
+ if exists(mask):
621
+ masked_out_value = orig_input
622
+
623
+ if self.return_zeros_for_masked_padding:
624
+ masked_out_value = torch.zeros_like(quantized_out)
625
+
626
+ quantized_out = einx.where(
627
+ 'b n, b n d, b n d -> b n d',
628
+ mask,
629
+ quantized_out,
630
+ masked_out_value
631
+ )
632
+
612
633
  # whether to early return the cross entropy loss
613
634
 
614
635
  if return_loss:
@@ -1154,6 +1154,7 @@ class VectorQuantize(Module):
1154
1154
 
1155
1155
  if need_transpose:
1156
1156
  x = rearrange(x, 'b d n -> b n d')
1157
+ orig_input = rearrange(orig_input, 'b d n -> b n d')
1157
1158
 
1158
1159
  # project input
1159
1160
 
@@ -1295,7 +1296,7 @@ class VectorQuantize(Module):
1295
1296
  commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
1296
1297
 
1297
1298
  if exists(mask):
1298
- commit_loss = einx.where('..., ... k, -> ... k', mask, commit_loss, 0.)
1299
+ commit_loss = einx.where('b n, b n ... k, -> b n ... k', mask, commit_loss, 0.)
1299
1300
 
1300
1301
  loss = commit_loss * self.commitment_weight if self.has_commitment_loss else commit_loss
1301
1302
 
@@ -1327,7 +1328,7 @@ class VectorQuantize(Module):
1327
1328
  commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
1328
1329
 
1329
1330
  if exists(mask):
1330
- commit_loss = einx.where('..., ... k, -> ... k', mask, commit_loss, 0.)
1331
+ commit_loss = einx.where('b n, b n ... k, -> b n ... k', mask, commit_loss, 0.)
1331
1332
 
1332
1333
  elif exists(mask):
1333
1334
  # with variable lengthed sequences
@@ -1386,20 +1387,6 @@ class VectorQuantize(Module):
1386
1387
 
1387
1388
  quantize = self.project_out(quantize)
1388
1389
 
1389
- # rearrange quantized embeddings
1390
-
1391
- if need_transpose:
1392
- quantize = rearrange(quantize, 'b n ... d -> b d n ...')
1393
-
1394
- if self.accept_image_fmap:
1395
- quantize = rearrange(quantize, 'b (h w) ... c -> b c h w ...', h = height, w = width)
1396
-
1397
- if self.accept_3d_fmap:
1398
- quantize = rearrange(quantize, "b (d h w) ... c -> b c d h w ...", d=depth, h=height, w=width)
1399
-
1400
- if only_one:
1401
- quantize = rearrange(quantize, 'b 1 ... d -> b ... d')
1402
-
1403
1390
  # if masking, only return quantized for where mask has True
1404
1391
 
1405
1392
  if exists(mask):
@@ -1408,8 +1395,11 @@ class VectorQuantize(Module):
1408
1395
  if self.return_zeros_for_masked_padding:
1409
1396
  masked_out_value = torch.zeros_like(orig_input)
1410
1397
 
1398
+ if is_topk:
1399
+ masked_out_value = repeat(masked_out_value, '... d -> ... k d', k = topk)
1400
+
1411
1401
  quantize = einx.where(
1412
- 'b n, b n ... d, b n d -> b n ... d',
1402
+ 'b n, b n ... d, b n ... d -> b n ... d',
1413
1403
  mask,
1414
1404
  quantize,
1415
1405
  masked_out_value
@@ -1422,6 +1412,20 @@ class VectorQuantize(Module):
1422
1412
  -1
1423
1413
  )
1424
1414
 
1415
+ # rearrange quantized embeddings
1416
+
1417
+ if need_transpose:
1418
+ quantize = rearrange(quantize, 'b n ... d -> b d n ...')
1419
+
1420
+ if self.accept_image_fmap:
1421
+ quantize = rearrange(quantize, 'b (h w) ... c -> b c h w ...', h = height, w = width)
1422
+
1423
+ if self.accept_3d_fmap:
1424
+ quantize = rearrange(quantize, "b (d h w) ... c -> b c d h w ...", d=depth, h=height, w=width)
1425
+
1426
+ if only_one:
1427
+ quantize = rearrange(quantize, 'b 1 ... d -> b ... d')
1428
+
1425
1429
  if not return_loss_breakdown:
1426
1430
  return quantize, embed_ind, loss
1427
1431