vector-quantize-pytorch 1.23.2__tar.gz → 1.23.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 (32) hide show
  1. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/vector_quantize_pytorch.py +8 -4
  4. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/.github/workflows/build.yml +0 -0
  5. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/.github/workflows/python-publish.yml +0 -0
  6. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/.github/workflows/test.yml +0 -0
  7. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/.gitignore +0 -0
  8. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/LICENSE +0 -0
  9. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/README.md +0 -0
  10. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/examples/autoencoder.py +0 -0
  11. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/examples/autoencoder_fsq.py +0 -0
  12. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/examples/autoencoder_lfq.py +0 -0
  13. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/examples/autoencoder_sim_vq.py +0 -0
  14. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/images/fsq.png +0 -0
  15. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/images/lfq.png +0 -0
  16. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/images/simvq.png +0 -0
  17. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/images/vq.png +0 -0
  18. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/ruff.toml +0 -0
  19. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/tests/test_latent_quantization.py +0 -0
  20. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/tests/test_lfq.py +0 -0
  21. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/tests/test_readme.py +0 -0
  22. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/__init__.py +0 -0
  23. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  24. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
  25. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  26. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  27. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
  28. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
  29. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  30. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/residual_vq.py +0 -0
  31. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/vector_quantize_pytorch/sim_vq.py +0 -0
  32. {vector_quantize_pytorch-1.23.2 → vector_quantize_pytorch-1.23.4}/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.23.2
3
+ Version: 1.23.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.23.2"
3
+ version = "1.23.4"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -273,12 +273,14 @@ def efficient_rotation_trick_transform(u, q, e):
273
273
  e = rearrange(e, 'b d -> b 1 d')
274
274
  w = l2norm(u + q, dim = 1).detach()
275
275
 
276
- return (
276
+ out = (
277
277
  e -
278
278
  2 * (e @ rearrange(w, 'b d -> b d 1') @ rearrange(w, 'b d -> b 1 d')) +
279
279
  2 * (e @ rearrange(u, 'b d -> b d 1').detach() @ rearrange(q, 'b d -> b 1 d').detach())
280
280
  )
281
281
 
282
+ return rearrange(out, '... 1 -> ...')
283
+
282
284
  def rotate_to(src, tgt):
283
285
  # rotation trick STE (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
284
286
  src, inverse = pack_one(src, '* d')
@@ -291,7 +293,7 @@ def rotate_to(src, tgt):
291
293
  safe_div(src, norm_src),
292
294
  safe_div(tgt, norm_tgt),
293
295
  src
294
- ).squeeze()
296
+ )
295
297
 
296
298
  rotated = rotated_tgt * safe_div(norm_tgt, norm_src).detach()
297
299
 
@@ -896,7 +898,7 @@ class VectorQuantize(Module):
896
898
  stochastic_sample_codes = False,
897
899
  sample_codebook_temp = 1.,
898
900
  straight_through = False,
899
- rotation_trick = True, # Propagate grads through VQ layer w/ rotation trick: https://arxiv.org/abs/2410.06424 by @cfifty
901
+ rotation_trick = None, # Propagate grads through VQ layer w/ rotation trick: https://arxiv.org/abs/2410.06424 by @cfifty
900
902
  sync_codebook = None,
901
903
  sync_affine_param = False,
902
904
  ema_update = True,
@@ -911,6 +913,8 @@ class VectorQuantize(Module):
911
913
  return_zeros_for_masked_padding = True
912
914
  ):
913
915
  super().__init__()
916
+ rotation_trick = default(rotation_trick, dim > 1) # only use rotation trick if feature dimension greater than 1
917
+
914
918
  self.dim = dim
915
919
  self.heads = heads
916
920
  self.separate_codebook_per_head = separate_codebook_per_head
@@ -1263,7 +1267,7 @@ class VectorQuantize(Module):
1263
1267
  # calculate codebook diversity loss (negative of entropy) if needed
1264
1268
 
1265
1269
  if self.has_codebook_diversity_loss:
1266
- prob = (-distances * self.codebook_diversity_temperature).softmax(dim = -1)
1270
+ prob = (distances * self.codebook_diversity_temperature).softmax(dim = -1)
1267
1271
  avg_prob = reduce(prob, '... n l -> n l', 'mean')
1268
1272
  codebook_diversity_loss = -entropy(avg_prob).mean()
1269
1273