vector-quantize-pytorch 1.28.0__tar.gz → 1.28.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 (18) hide show
  1. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/PKG-INFO +13 -1
  2. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/README.md +12 -0
  3. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/pyproject.toml +1 -1
  4. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/finite_scalar_quantization.py +22 -1
  5. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/lookup_free_quantization.py +22 -1
  6. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +2 -1
  7. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/.gitignore +0 -0
  8. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/LICENSE +0 -0
  9. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/__init__.py +0 -0
  10. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/binary_mapper.py +0 -0
  11. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
  12. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  13. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
  14. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
  15. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  16. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_vq.py +0 -0
  17. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/sim_vq.py +0 -0
  18. {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/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.28.0
3
+ Version: 1.28.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
@@ -859,3 +859,15 @@ assert loss.item() >= 0
859
859
  url = {https://arxiv.org/abs/2509.10140},
860
860
  }
861
861
  ```
862
+
863
+ ```bibtex
864
+ @misc{chee2024quip2bitquantizationlarge,
865
+ title = {QuIP: 2-Bit Quantization of Large Language Models With Guarantees},
866
+ author = {Jerry Chee and Yaohui Cai and Volodymyr Kuleshov and Christopher De Sa},
867
+ year = {2024},
868
+ eprint = {2307.13304},
869
+ archivePrefix = {arXiv},
870
+ primaryClass = {cs.LG},
871
+ url = {https://arxiv.org/abs/2307.13304},
872
+ }
873
+ ```
@@ -815,3 +815,15 @@ assert loss.item() >= 0
815
815
  url = {https://arxiv.org/abs/2509.10140},
816
816
  }
817
817
  ```
818
+
819
+ ```bibtex
820
+ @misc{chee2024quip2bitquantizationlarge,
821
+ title = {QuIP: 2-Bit Quantization of Large Language Models With Guarantees},
822
+ author = {Jerry Chee and Yaohui Cai and Volodymyr Kuleshov and Christopher De Sa},
823
+ year = {2024},
824
+ eprint = {2307.13304},
825
+ archivePrefix = {arXiv},
826
+ primaryClass = {cs.LG},
827
+ url = {https://arxiv.org/abs/2307.13304},
828
+ }
829
+ ```
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.28.0"
3
+ version = "1.28.2"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -76,7 +76,8 @@ class FSQ(Module):
76
76
  force_quantization_f32 = True,
77
77
  preserve_symmetry = False,
78
78
  noise_dropout = 0.,
79
- bound_hard_clamp = False # for residual fsq, if input is pre-softclamped to the right range
79
+ bound_hard_clamp = False, # for residual fsq, if input is pre-softclamped to the right range
80
+ orthogonal_rotation = False # increase codebook utilization. ensure levels are symmetric! https://arxiv.org/abs/2307.13304v2
80
81
  ):
81
82
  super().__init__()
82
83
 
@@ -132,6 +133,17 @@ class FSQ(Module):
132
133
 
133
134
  self.bound_hard_clamp = bound_hard_clamp
134
135
 
136
+ self.orthogonal_rotation = orthogonal_rotation
137
+
138
+ if orthogonal_rotation:
139
+ is_symmetric = len(set(levels)) == 1
140
+ if not is_symmetric:
141
+ print('orthogonal_rotation is not recommended for FSQ with asymmetric levels (i.e. where the number of bins differ across dimensions)')
142
+
143
+ orthogonal_rot = torch.empty(codebook_dim, codebook_dim)
144
+ nn.init.orthogonal_(orthogonal_rot)
145
+ self.register_buffer('orthogonal_rot', orthogonal_rot)
146
+
135
147
  def bound(self, z, eps = 1e-3, hard_clamp = False):
136
148
  """ Bound `z`, an array of shape (..., d). """
137
149
  maybe_tanh = tanh if not hard_clamp else partial(clamp, min = -1., max = 1.)
@@ -219,6 +231,9 @@ class FSQ(Module):
219
231
 
220
232
  codes = self._indices_to_codes(indices)
221
233
 
234
+ if self.orthogonal_rotation:
235
+ codes = codes @ self.orthogonal_rot.t()
236
+
222
237
  if self.keep_num_codebooks_dim:
223
238
  codes = rearrange(codes, '... c d -> ... (c d)')
224
239
 
@@ -253,6 +268,9 @@ class FSQ(Module):
253
268
 
254
269
  z = rearrange(z, 'b n (c d) -> b n c d', c = self.num_codebooks)
255
270
 
271
+ if self.orthogonal_rotation:
272
+ z = z @ self.orthogonal_rot
273
+
256
274
  # whether to force quantization step to be full precision or not
257
275
 
258
276
  force_f32 = self.force_quantization_f32
@@ -275,6 +293,9 @@ class FSQ(Module):
275
293
 
276
294
  codes = self.maybe_apply_noise(codes)
277
295
 
296
+ if self.orthogonal_rotation:
297
+ codes = codes @ self.orthogonal_rot.t()
298
+
278
299
  codes = rearrange(codes, 'b n c d -> b n (c d)')
279
300
 
280
301
  codes = codes.to(orig_dtype)
@@ -116,7 +116,8 @@ class LFQ(Module):
116
116
  experimental_softplus_entropy_loss = False,
117
117
  entropy_loss_offset = 5., # how much to shift the loss before softplus
118
118
  spherical = False, # from https://arxiv.org/abs/2406.07548
119
- force_quantization_f32 = True # will force the quantization step to be full precision
119
+ force_quantization_f32 = True, # will force the quantization step to be full precision
120
+ orthogonal_rotation = False # increase codebook utilization without aux losses, inspired by https://arxiv.org/abs/2307.13304v2
120
121
  ):
121
122
  super().__init__()
122
123
 
@@ -165,6 +166,15 @@ class LFQ(Module):
165
166
  self.spherical = spherical
166
167
  self.maybe_l2norm = (lambda t: l2norm(t) * self.codebook_scale) if spherical else identity
167
168
 
169
+ # orthogonal rotation
170
+
171
+ self.orthogonal_rotation = orthogonal_rotation
172
+
173
+ if orthogonal_rotation:
174
+ orthogonal_rot = torch.empty(codebook_dim, codebook_dim)
175
+ nn.init.orthogonal_(orthogonal_rot)
176
+ self.register_buffer('orthogonal_rot', orthogonal_rot)
177
+
168
178
  # entropy aux loss related weights
169
179
 
170
180
  assert 0 < frac_per_sample_entropy <= 1.
@@ -234,6 +244,9 @@ class LFQ(Module):
234
244
 
235
245
  codes = self.maybe_l2norm(codes)
236
246
 
247
+ if self.orthogonal_rotation:
248
+ codes = codes @ self.orthogonal_rot.t()
249
+
237
250
  codes = rearrange(codes, '... c d -> ... (c d)')
238
251
 
239
252
  # whether to project codes out to original dimensions
@@ -287,6 +300,9 @@ class LFQ(Module):
287
300
 
288
301
  x = rearrange(x, 'b n (c d) -> b n c d', c = self.num_codebooks)
289
302
 
303
+ if self.orthogonal_rotation:
304
+ x = x @ self.orthogonal_rot
305
+
290
306
  # maybe l2norm
291
307
 
292
308
  x = self.maybe_l2norm(x)
@@ -412,6 +428,11 @@ class LFQ(Module):
412
428
  if force_f32:
413
429
  x = x.type(orig_dtype)
414
430
 
431
+ # rotate back if needed
432
+
433
+ if self.orthogonal_rotation:
434
+ x = x @ self.orthogonal_rot.t()
435
+
415
436
  # merge back codebook dim
416
437
 
417
438
  x = rearrange(x, 'b n c d -> b n (c d)')
@@ -322,7 +322,7 @@ def directional_reparam(src, tgt, noise_variance = 5e-3):
322
322
  error_dir_norm = error_dir.norm(dim = -1, keepdim = True)
323
323
 
324
324
  noised_dir = error_dir + sqrt(noise_variance) * torch.randn_like(error_dir)
325
- unit_noised_dir = l2norm(noised_dir)
325
+ unit_noised_dir = l2norm(noised_dir).detach()
326
326
 
327
327
  return src + unit_noised_dir * error_dir_norm
328
328
 
@@ -860,6 +860,7 @@ class VectorQuantize(Module):
860
860
  assert at_most_one_of(straight_through, rotation_trick, directional_reparam)
861
861
  self.rotation_trick = rotation_trick
862
862
 
863
+ assert not (directional_reparam and threshold_ema_dead_code == 0), 'periodic dead code replacement should be enabled when differential reparam method is turned on'
863
864
  self.directional_reparam = directional_reparam
864
865
  self.directional_reparam_variance = directional_reparam_variance
865
866