vector-quantize-pytorch 1.20.2__tar.gz → 1.20.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 (29) hide show
  1. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_sim_vq.py +10 -2
  3. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/pyproject.toml +1 -1
  4. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/tests/test_readme.py +14 -0
  5. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/sim_vq.py +23 -5
  6. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/vector_quantize_pytorch.py +9 -9
  7. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/build.yml +0 -0
  8. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/python-publish.yml +0 -0
  9. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/test.yml +0 -0
  10. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.gitignore +0 -0
  11. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/LICENSE +0 -0
  12. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/README.md +0 -0
  13. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder.py +0 -0
  14. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_fsq.py +0 -0
  15. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_lfq.py +0 -0
  16. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/images/fsq.png +0 -0
  17. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/images/lfq.png +0 -0
  18. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/images/vq.png +0 -0
  19. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/ruff.toml +0 -0
  20. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/tests/test_latent_quantization.py +0 -0
  21. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/__init__.py +0 -0
  22. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  23. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
  24. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  25. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  26. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
  27. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
  28. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/residual_vq.py +0 -0
  29. {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.20.2
3
+ Version: 1.20.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
@@ -14,7 +14,10 @@ lr = 3e-4
14
14
  train_iter = 10000
15
15
  num_codes = 256
16
16
  seed = 1234
17
- rotation_trick = True
17
+
18
+ rotation_trick = True # rotation trick instead ot straight-through
19
+ use_mlp = True # use a one layer mlp with relu instead of linear
20
+
18
21
  device = "cuda" if torch.cuda.is_available() else "cpu"
19
22
 
20
23
  def SimVQAutoEncoder(**vq_kwargs):
@@ -77,7 +80,12 @@ torch.random.manual_seed(seed)
77
80
 
78
81
  model = SimVQAutoEncoder(
79
82
  codebook_size = num_codes,
80
- rotation_trick = rotation_trick
83
+ rotation_trick = rotation_trick,
84
+ codebook_transform = nn.Sequential(
85
+ nn.Linear(32, 128),
86
+ nn.ReLU(),
87
+ nn.Linear(128, 32),
88
+ ) if use_mlp else None
81
89
  ).to(device)
82
90
 
83
91
  opt = torch.optim.AdamW(model.parameters(), lr=lr)
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.20.2"
3
+ version = "1.20.4"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -362,3 +362,17 @@ def test_latent_q():
362
362
 
363
363
  assert image_feats.shape == quantized.shape
364
364
  assert (quantized == quantizer.indices_to_codes(indices)).all()
365
+
366
+ def test_sim_vq():
367
+ from vector_quantize_pytorch import SimVQ
368
+
369
+ sim_vq = SimVQ(
370
+ dim = 512,
371
+ codebook_size = 1024,
372
+ )
373
+
374
+ x = torch.randn(1, 1024, 512)
375
+ quantized, indices, commit_loss = sim_vq(x)
376
+
377
+ assert x.shape == quantized.shape
378
+ assert torch.allclose(quantized, sim_vq.indices_to_codes(indices), atol = 1e-6)
@@ -9,7 +9,7 @@ import torch.nn.functional as F
9
9
  from einx import get_at
10
10
  from einops import rearrange, pack, unpack
11
11
 
12
- from vector_quantize_pytorch.vector_quantize_pytorch import rotate_from_to
12
+ from vector_quantize_pytorch.vector_quantize_pytorch import rotate_to
13
13
 
14
14
  # helper functions
15
15
 
@@ -56,9 +56,9 @@ class SimVQ(Module):
56
56
  if not exists(codebook_transform):
57
57
  codebook_transform = nn.Linear(dim, dim, bias = False)
58
58
 
59
- self.codebook_to_codes = codebook_transform
59
+ self.code_transform = codebook_transform
60
60
 
61
- self.register_buffer('codebook', codebook)
61
+ self.register_buffer('frozen_codebook', codebook)
62
62
 
63
63
 
64
64
  # whether to use rotation trick from Fifty et al.
@@ -70,6 +70,24 @@ class SimVQ(Module):
70
70
 
71
71
  self.input_to_quantize_commit_loss_weight = input_to_quantize_commit_loss_weight
72
72
 
73
+ @property
74
+ def codebook(self):
75
+ return self.code_transform(self.frozen_codebook)
76
+
77
+ def indices_to_codes(
78
+ self,
79
+ indices
80
+ ):
81
+ implicit_codebook = self.codebook
82
+
83
+ frozen_codes = get_at('[c] d, b ... -> b ... d', self.frozen_codebook, indices)
84
+ quantized = self.code_transform(frozen_codes)
85
+
86
+ if self.accept_image_fmap:
87
+ quantized = rearrange(quantized, 'b ... d -> b d ...')
88
+
89
+ return quantized
90
+
73
91
  def forward(
74
92
  self,
75
93
  x
@@ -78,7 +96,7 @@ class SimVQ(Module):
78
96
  x = rearrange(x, 'b d h w -> b h w d')
79
97
  x, inverse_pack = pack_one(x, 'b * d')
80
98
 
81
- implicit_codebook = self.codebook_to_codes(self.codebook)
99
+ implicit_codebook = self.codebook
82
100
 
83
101
  with torch.no_grad():
84
102
  dist = torch.cdist(x, implicit_codebook)
@@ -94,7 +112,7 @@ class SimVQ(Module):
94
112
 
95
113
  if self.rotation_trick:
96
114
  # rotation trick from @cfifty
97
- quantized = rotate_from_to(quantized, x)
115
+ quantized = rotate_to(x, quantized)
98
116
  else:
99
117
 
100
118
  commit_loss = (
@@ -250,21 +250,21 @@ def efficient_rotation_trick_transform(u, q, e):
250
250
  2 * (e @ rearrange(u, 'b d -> b d 1').detach() @ rearrange(q, 'b d -> b 1 d').detach())
251
251
  )
252
252
 
253
- def rotate_from_to(src, tgt):
253
+ def rotate_to(src, tgt):
254
254
  # rotation trick STE (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
255
- tgt, inverse = pack_one(tgt, '* d')
256
- src, _ = pack_one(src, '* d')
255
+ src, inverse = pack_one(src, '* d')
256
+ tgt, _ = pack_one(tgt, '* d')
257
257
 
258
- norm_tgt = tgt.norm(dim = -1, keepdim = True)
259
258
  norm_src = src.norm(dim = -1, keepdim = True)
259
+ norm_tgt = tgt.norm(dim = -1, keepdim = True)
260
260
 
261
- rotated_src = efficient_rotation_trick_transform(
262
- safe_div(tgt, norm_tgt),
261
+ rotated_tgt = efficient_rotation_trick_transform(
263
262
  safe_div(src, norm_src),
264
- tgt
263
+ safe_div(tgt, norm_tgt),
264
+ src
265
265
  ).squeeze()
266
266
 
267
- rotated = rotated_src * safe_div(norm_src, norm_tgt).detach()
267
+ rotated = rotated_tgt * safe_div(norm_tgt, norm_src).detach()
268
268
 
269
269
  return inverse(rotated)
270
270
 
@@ -1118,7 +1118,7 @@ class VectorQuantize(Module):
1118
1118
  commit_quantize = maybe_detach(quantize)
1119
1119
 
1120
1120
  if self.rotation_trick:
1121
- quantize = rotate_from_to(quantize, x)
1121
+ quantize = rotate_to(x, quantize)
1122
1122
  else:
1123
1123
  # standard STE to get gradients through VQ layer.
1124
1124
  quantize = x + (quantize - x).detach()