vector-quantize-pytorch 1.22.6__tar.gz → 1.22.7__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.22.6 → vector_quantize_pytorch-1.22.7}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/finite_scalar_quantization.py +8 -25
  4. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/.github/workflows/build.yml +0 -0
  5. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/.github/workflows/python-publish.yml +0 -0
  6. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/.github/workflows/test.yml +0 -0
  7. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/.gitignore +0 -0
  8. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/LICENSE +0 -0
  9. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/README.md +0 -0
  10. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/examples/autoencoder.py +0 -0
  11. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/examples/autoencoder_fsq.py +0 -0
  12. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/examples/autoencoder_lfq.py +0 -0
  13. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/examples/autoencoder_sim_vq.py +0 -0
  14. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/images/fsq.png +0 -0
  15. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/images/lfq.png +0 -0
  16. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/images/simvq.png +0 -0
  17. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/images/vq.png +0 -0
  18. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/ruff.toml +0 -0
  19. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/tests/test_latent_quantization.py +0 -0
  20. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/tests/test_lfq.py +0 -0
  21. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/tests/test_readme.py +0 -0
  22. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/__init__.py +0 -0
  23. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/latent_quantization.py +0 -0
  24. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  25. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  26. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/residual_fsq.py +0 -0
  27. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/residual_lfq.py +0 -0
  28. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  29. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/residual_vq.py +0 -0
  30. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/sim_vq.py +0 -0
  31. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/utils.py +0 -0
  32. {vector_quantize_pytorch-1.22.6 → vector_quantize_pytorch-1.22.7}/vector_quantize_pytorch/vector_quantize_pytorch.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.22.6
3
+ Version: 1.22.7
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.22.6"
3
+ version = "1.22.7"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -137,38 +137,21 @@ class FSQ(Module):
137
137
 
138
138
  def quantize(self, z):
139
139
  """ Quantizes z, returns quantized zhat, same shape as z. """
140
+ shape, device, noise_dropout, preserve_symmetry, half_width = z.shape[0], z.device, self.noise_dropout, self.preserve_symmetry, (self._levels // 2)
140
141
 
141
- preserve_symmetry = self.preserve_symmetry
142
- half_width = self._levels // 2
142
+ # determine where to add a random offset elementwise
143
+ # if using noise dropout
144
+
145
+ if self.training and noise_dropout > 0.:
146
+ offset_mask = torch.bernoulli(torch.full_like(z, noise_dropout)).bool()
147
+ offset = (torch.rand_like(z) - 0.5) / half_width
148
+ z = torch.where(offset_mask, z + offset, z)
143
149
 
144
150
  if preserve_symmetry:
145
151
  quantized = round_ste(self.symmetry_preserving_bound(z)) / half_width
146
152
  else:
147
153
  quantized = round_ste(self.bound(z)) / half_width
148
154
 
149
- if not self.training:
150
- return quantized
151
-
152
- batch, device, noise_dropout = z.shape[0], z.device, self.noise_dropout
153
- unquantized = z
154
-
155
- # determine where to quantize elementwise
156
-
157
- quantize_mask = torch.bernoulli(
158
- torch.full((batch,), noise_dropout, device = device)
159
- ).bool()
160
-
161
- quantized = einx.where('b, b ..., b ...', quantize_mask, unquantized, quantized)
162
-
163
- # determine where to add a random offset elementwise
164
-
165
- offset_mask = torch.bernoulli(
166
- torch.full((batch,), noise_dropout, device = device)
167
- ).bool()
168
-
169
- offset = (torch.rand_like(z) - 0.5) / half_width
170
- quantized = einx.where('b, b ..., b ...', offset_mask, unquantized + offset, quantized)
171
-
172
155
  return quantized
173
156
 
174
157
  def _scale_and_shift(self, zhat_normalized):