vector-quantize-pytorch 1.18.0__tar.gz → 1.18.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 (26) hide show
  1. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/tests/test_readme.py +5 -2
  4. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/lookup_free_quantization.py +1 -1
  5. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +37 -26
  6. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/build.yml +0 -0
  7. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/python-publish.yml +0 -0
  8. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/test.yml +0 -0
  9. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.gitignore +0 -0
  10. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/LICENSE +0 -0
  11. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/README.md +0 -0
  12. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder.py +0 -0
  13. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder_fsq.py +0 -0
  14. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder_lfq.py +0 -0
  15. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/images/fsq.png +0 -0
  16. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/images/lfq.png +0 -0
  17. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/images/vq.png +0 -0
  18. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/ruff.toml +0 -0
  19. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/tests/test_latent_quantization.py +0 -0
  20. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/__init__.py +0 -0
  21. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  22. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
  23. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  24. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
  25. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
  26. {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/residual_vq.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.18.0
3
+ Version: 1.18.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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.18.0"
3
+ version = "1.18.2"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -5,8 +5,10 @@ def exists(v):
5
5
  return v is not None
6
6
 
7
7
  @pytest.mark.parametrize('use_cosine_sim', (True, False))
8
+ @pytest.mark.parametrize('rotation_trick', (True, False))
8
9
  def test_vq(
9
- use_cosine_sim
10
+ use_cosine_sim,
11
+ rotation_trick
10
12
  ):
11
13
  from vector_quantize_pytorch import VectorQuantize
12
14
 
@@ -15,7 +17,8 @@ def test_vq(
15
17
  codebook_size = 512, # codebook size
16
18
  decay = 0.8, # the exponential moving average decay, lower means the dictionary will change faster
17
19
  commitment_weight = 1., # the weight on the commitment loss
18
- use_cosine_sim = use_cosine_sim
20
+ use_cosine_sim = use_cosine_sim,
21
+ rotation_trick = rotation_trick
19
22
  )
20
23
 
21
24
  x = torch.randn(1, 1024, 256)
@@ -100,7 +100,7 @@ class LFQ(Module):
100
100
  dim = None,
101
101
  codebook_size = None,
102
102
  entropy_loss_weight = 0.1,
103
- commitment_loss_weight = 0.25,
103
+ commitment_loss_weight = 0.,
104
104
  diversity_gamma = 1.,
105
105
  straight_through_activation = nn.Identity(),
106
106
  num_codebooks = 1,
@@ -28,8 +28,11 @@ def noop(*args, **kwargs):
28
28
  def identity(t):
29
29
  return t
30
30
 
31
- def l2norm(t):
32
- return F.normalize(t, p = 2, dim = -1)
31
+ def l2norm(t, dim = -1, eps = 1e-6):
32
+ return F.normalize(t, p = 2, dim = dim, eps = eps)
33
+
34
+ def safe_div(num, den, eps = 1e-6):
35
+ return num / den.clamp(min = eps)
33
36
 
34
37
  def Sequential(*modules):
35
38
  modules = [*filter(exists, modules)]
@@ -73,6 +76,19 @@ def lens_to_mask(lens, max_length):
73
76
  seq = torch.arange(max_length, device = lens.device)
74
77
  return seq < lens[:, None]
75
78
 
79
+ def efficient_rotation_trick_transform(u, q, e):
80
+ """
81
+ 4.2 in https://arxiv.org/abs/2410.06424
82
+ """
83
+ e = rearrange(e, 'b d -> b 1 d')
84
+ w = l2norm(u + q, dim = 1).detach()
85
+
86
+ return (
87
+ e -
88
+ 2 * (e @ rearrange(w, 'b d -> b d 1') @ rearrange(w, 'b d -> b 1 d')) +
89
+ 2 * (e @ rearrange(u, 'b d -> b d 1').detach() @ rearrange(q, 'b d -> b 1 d').detach())
90
+ )
91
+
76
92
  def uniform_init(*shape):
77
93
  t = torch.empty(shape)
78
94
  nn.init.kaiming_uniform_(t)
@@ -811,7 +827,7 @@ class VectorQuantize(Module):
811
827
  stochastic_sample_codes = False,
812
828
  sample_codebook_temp = 1.,
813
829
  straight_through = False,
814
- rotation_trick = True, # Propagate grads through VQ layer w/ rotation trick: https://arxiv.org/abs/2410.06424.
830
+ rotation_trick = True, # Propagate grads through VQ layer w/ rotation trick: https://arxiv.org/abs/2410.06424 by @cfifty
815
831
  reinmax = False, # using reinmax for improved straight-through, assuming straight through helps at all
816
832
  sync_codebook = None,
817
833
  sync_affine_param = False,
@@ -946,13 +962,6 @@ class VectorQuantize(Module):
946
962
 
947
963
  self._codebook.embed.copy_(codes)
948
964
 
949
- @staticmethod
950
- def rotation_trick_transform(u, q, e):
951
- w = ((u + q) / torch.norm(u + q, dim=1, keepdim=True)).detach()
952
- e = e - 2 * torch.bmm(torch.bmm(e, w.unsqueeze(-1)), w.unsqueeze(1)) + 2 * torch.bmm(
953
- torch.bmm(e, u.unsqueeze(-1).detach()), q.unsqueeze(1).detach())
954
- return e
955
-
956
965
  def get_codes_from_indices(self, indices):
957
966
  codebook = self.codebook
958
967
  is_multiheaded = codebook.ndim > 2
@@ -1103,23 +1112,25 @@ class VectorQuantize(Module):
1103
1112
 
1104
1113
  commit_quantize = maybe_detach(quantize)
1105
1114
 
1106
- # Use the rotation trick (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
1107
1115
  if self.rotation_trick:
1108
- init_shape = x.shape
1109
- x = x.reshape(-1, init_shape[-1])
1110
- quantize = quantize.reshape(-1, init_shape[-1])
1111
-
1112
- eps = 1e-6 # For numerical stability if any vector is close to 0 norm.
1113
- rot_quantize = self.rotation_trick_transform(
1114
- x / (torch.norm(x, dim=1, keepdim=True) + eps),
1115
- quantize / (torch.norm(quantize, dim=1, keepdim=True) + eps),
1116
- x.unsqueeze(1)).squeeze()
1117
- quantize = rot_quantize * (torch.norm(quantize, dim=1, keepdim=True)
1118
- / (torch.norm(x, dim=1, keepdim=True) + 1e-6)).detach()
1119
-
1120
- x = x.reshape(init_shape)
1121
- quantize = quantize.reshape(init_shape)
1122
- else: # Use STE to get gradients through VQ layer.
1116
+ # rotation trick STE (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
1117
+ x, inverse = pack_one(x, '* d')
1118
+ quantize, _ = pack_one(quantize, '* d')
1119
+
1120
+ norm_x = x.norm(dim = -1, keepdim = True)
1121
+ norm_quantize = quantize.norm(dim = -1, keepdim = True)
1122
+
1123
+ rot_quantize = efficient_rotation_trick_transform(
1124
+ safe_div(x, norm_x),
1125
+ safe_div(quantize, norm_quantize),
1126
+ x
1127
+ ).squeeze()
1128
+
1129
+ quantize = rot_quantize * safe_div(norm_quantize, norm_x).detach()
1130
+
1131
+ x, quantize = inverse(x), inverse(quantize)
1132
+ else:
1133
+ # standard STE to get gradients through VQ layer.
1123
1134
  quantize = x + (quantize - x).detach()
1124
1135
 
1125
1136
  if self.sync_update_v > 0.: