vector-quantize-pytorch 1.18.2__tar.gz → 1.18.5__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.2 → vector_quantize_pytorch-1.18.5}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/residual_fsq.py +25 -2
  4. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/.github/workflows/build.yml +0 -0
  5. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/.github/workflows/python-publish.yml +0 -0
  6. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/.github/workflows/test.yml +0 -0
  7. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/.gitignore +0 -0
  8. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/LICENSE +0 -0
  9. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/README.md +0 -0
  10. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/examples/autoencoder.py +0 -0
  11. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/examples/autoencoder_fsq.py +0 -0
  12. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/examples/autoencoder_lfq.py +0 -0
  13. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/images/fsq.png +0 -0
  14. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/images/lfq.png +0 -0
  15. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/images/vq.png +0 -0
  16. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/ruff.toml +0 -0
  17. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/tests/test_latent_quantization.py +0 -0
  18. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/tests/test_readme.py +0 -0
  19. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/__init__.py +0 -0
  20. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  21. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/latent_quantization.py +0 -0
  22. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  23. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  24. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/residual_lfq.py +0 -0
  25. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/residual_vq.py +0 -0
  26. {vector_quantize_pytorch-1.18.2 → vector_quantize_pytorch-1.18.5}/vector_quantize_pytorch/vector_quantize_pytorch.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.18.2
3
+ Version: 1.18.5
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.2"
3
+ version = "1.18.5"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -38,9 +38,10 @@ class ResidualFSQ(Module):
38
38
  def __init__(
39
39
  self,
40
40
  *,
41
- dim,
42
41
  levels: List[int],
43
42
  num_quantizers,
43
+ dim = None,
44
+ is_channel_first = False,
44
45
  quantize_dropout = False,
45
46
  quantize_dropout_cutoff_index = 0,
46
47
  quantize_dropout_multiple_of = 1,
@@ -48,12 +49,14 @@ class ResidualFSQ(Module):
48
49
  ):
49
50
  super().__init__()
50
51
  codebook_dim = len(levels)
52
+ dim = default(dim, codebook_dim)
51
53
 
52
54
  requires_projection = codebook_dim != dim
53
55
  self.project_in = nn.Linear(dim, codebook_dim) if requires_projection else nn.Identity()
54
56
  self.project_out = nn.Linear(codebook_dim, dim) if requires_projection else nn.Identity()
55
57
  self.has_projections = requires_projection
56
58
 
59
+ self.is_channel_first = is_channel_first
57
60
  self.num_quantizers = num_quantizers
58
61
 
59
62
  self.levels = levels
@@ -143,10 +146,18 @@ class ResidualFSQ(Module):
143
146
  ):
144
147
  num_quant, quant_dropout_multiple_of, device = self.num_quantizers, self.quantize_dropout_multiple_of, x.device
145
148
 
149
+ # handle channel first
150
+
151
+ if self.is_channel_first:
152
+ x = rearrange(x, 'b d ... -> b ... d')
153
+ x, ps = pack([x], 'b * d')
154
+
155
+ # maybe project in
156
+
146
157
  x = self.project_in(x)
147
158
 
148
159
  quantized_out = 0.
149
- residual = first(self.layers).bound(x)
160
+ residual = x
150
161
 
151
162
  all_indices = []
152
163
 
@@ -175,6 +186,7 @@ class ResidualFSQ(Module):
175
186
  continue
176
187
 
177
188
  quantized, indices = layer(residual / scale)
189
+
178
190
  quantized = quantized * scale
179
191
 
180
192
  residual = residual - quantized.detach()
@@ -190,6 +202,17 @@ class ResidualFSQ(Module):
190
202
 
191
203
  all_indices = torch.stack(all_indices, dim = -1)
192
204
 
205
+ # channel first out
206
+
207
+ if self.is_channel_first:
208
+ quantized_out, = unpack(quantized_out, ps, 'b * d')
209
+ all_indices, = unpack(all_indices, ps, 'b * d')
210
+
211
+ quantized_out = rearrange(quantized_out, 'b ... d -> b d ...')
212
+ all_indices = rearrange(all_indices, 'b ... d -> b d ...')
213
+
214
+ # return
215
+
193
216
  ret = (quantized_out, all_indices)
194
217
 
195
218
  if not return_all_codes: