vector-quantize-pytorch 1.28.2__tar.gz → 1.28.3__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.2 → vector_quantize_pytorch-1.28.3}/PKG-INFO +16 -3
  2. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/README.md +15 -2
  3. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/pyproject.toml +1 -1
  4. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/vector_quantize_pytorch.py +35 -10
  5. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/.gitignore +0 -0
  6. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/LICENSE +0 -0
  7. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/__init__.py +0 -0
  8. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/binary_mapper.py +0 -0
  9. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  10. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/latent_quantization.py +0 -0
  11. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  12. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  13. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_fsq.py +0 -0
  14. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_lfq.py +0 -0
  15. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  16. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_vq.py +0 -0
  17. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/sim_vq.py +0 -0
  18. {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/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.2
3
+ Version: 1.28.3
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
@@ -180,6 +180,19 @@ vq_layer = VectorQuantize(
180
180
  )
181
181
  ```
182
182
 
183
+ Alternatively, the <a href="https://openreview.net/forum?id=KRVnpTbx7R">DiVeQ paper</a> proposes to model quantization as the addition of a simulated quantization error to the input vector. The direction of this simulated error is aligned with the nearest codeword (if ```directional_reparam_variance``` is chosen small), and its magnitude equals the actual quantization error magnitude. In this way, the quantized output becomes a differentiable function of both input vector and selected codeword, and thus, DiVeQ provides valid gradients to learn the codebook without requiring any auxiliary losses. You can enable or disable this feature with ```directional_reparam = True/False``` in the ```VectorQuantize``` class.
184
+
185
+ ```python
186
+ from vector_quantize_pytorch import VectorQuantize
187
+
188
+ vq_layer = VectorQuantize(
189
+ dim = 256,
190
+ codebook_size = 256,
191
+ directional_reparam = True, # Set to False to use the STE gradient estimator or True to use the DiVeQ method.
192
+ directional_reparam_variance = 5e-3
193
+ )
194
+ ```
195
+
183
196
  ## Increasing codebook usage
184
197
 
185
198
  This repository will contain a few techniques from various papers to combat "dead" codebook entries, which is a common problem when using vector quantizers.
@@ -862,12 +875,12 @@ assert loss.item() >= 0
862
875
 
863
876
  ```bibtex
864
877
  @misc{chee2024quip2bitquantizationlarge,
865
- title = {QuIP: 2-Bit Quantization of Large Language Models With Guarantees},
878
+ title = {QuIP: 2-Bit Quantization of Large Language Models With Guarantees},
866
879
  author = {Jerry Chee and Yaohui Cai and Volodymyr Kuleshov and Christopher De Sa},
867
880
  year = {2024},
868
881
  eprint = {2307.13304},
869
882
  archivePrefix = {arXiv},
870
883
  primaryClass = {cs.LG},
871
- url = {https://arxiv.org/abs/2307.13304},
884
+ url = {https://arxiv.org/abs/2307.13304},
872
885
  }
873
886
  ```
@@ -136,6 +136,19 @@ vq_layer = VectorQuantize(
136
136
  )
137
137
  ```
138
138
 
139
+ Alternatively, the <a href="https://openreview.net/forum?id=KRVnpTbx7R">DiVeQ paper</a> proposes to model quantization as the addition of a simulated quantization error to the input vector. The direction of this simulated error is aligned with the nearest codeword (if ```directional_reparam_variance``` is chosen small), and its magnitude equals the actual quantization error magnitude. In this way, the quantized output becomes a differentiable function of both input vector and selected codeword, and thus, DiVeQ provides valid gradients to learn the codebook without requiring any auxiliary losses. You can enable or disable this feature with ```directional_reparam = True/False``` in the ```VectorQuantize``` class.
140
+
141
+ ```python
142
+ from vector_quantize_pytorch import VectorQuantize
143
+
144
+ vq_layer = VectorQuantize(
145
+ dim = 256,
146
+ codebook_size = 256,
147
+ directional_reparam = True, # Set to False to use the STE gradient estimator or True to use the DiVeQ method.
148
+ directional_reparam_variance = 5e-3
149
+ )
150
+ ```
151
+
139
152
  ## Increasing codebook usage
140
153
 
141
154
  This repository will contain a few techniques from various papers to combat "dead" codebook entries, which is a common problem when using vector quantizers.
@@ -818,12 +831,12 @@ assert loss.item() >= 0
818
831
 
819
832
  ```bibtex
820
833
  @misc{chee2024quip2bitquantizationlarge,
821
- title = {QuIP: 2-Bit Quantization of Large Language Models With Guarantees},
834
+ title = {QuIP: 2-Bit Quantization of Large Language Models With Guarantees},
822
835
  author = {Jerry Chee and Yaohui Cai and Volodymyr Kuleshov and Christopher De Sa},
823
836
  year = {2024},
824
837
  eprint = {2307.13304},
825
838
  archivePrefix = {arXiv},
826
839
  primaryClass = {cs.LG},
827
- url = {https://arxiv.org/abs/2307.13304},
840
+ url = {https://arxiv.org/abs/2307.13304},
828
841
  }
829
842
  ```
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.28.2"
3
+ version = "1.28.3"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -314,7 +314,7 @@ def rotate_to(src, tgt):
314
314
 
315
315
  return inverse(rotated)
316
316
 
317
- # directional reparam related
317
+ # directional reparameterization (DiVeQ) method to learn the codebook
318
318
  # figure 1. https://openreview.net/forum?id=KRVnpTbx7R
319
319
 
320
320
  def directional_reparam(src, tgt, noise_variance = 5e-3):
@@ -389,7 +389,10 @@ class Codebook(Module):
389
389
 
390
390
  self.kmeans_iters = kmeans_iters
391
391
  self.eps = eps
392
+
392
393
  self.threshold_ema_dead_code = threshold_ema_dead_code
394
+ self.has_dead_code_replacement = threshold_ema_dead_code > 0
395
+
393
396
  self.reset_cluster_size = default(reset_cluster_size, threshold_ema_dead_code)
394
397
 
395
398
  assert callable(gumbel_sample)
@@ -550,7 +553,7 @@ class Codebook(Module):
550
553
  self.embed_avg.data[ind][mask] = sampled * self.reset_cluster_size
551
554
 
552
555
  def expire_codes_(self, batch_samples):
553
- if self.threshold_ema_dead_code == 0:
556
+ if not self.has_dead_code_replacement or not self.training:
554
557
  return
555
558
 
556
559
  expired_codes = self.cluster_size < self.threshold_ema_dead_code
@@ -571,7 +574,7 @@ class Codebook(Module):
571
574
 
572
575
  self.embed.data.copy_(embed_normalized)
573
576
 
574
- def update_ema_part(
577
+ def track_cluster_size_and_embed_avg(
575
578
  self,
576
579
  flatten,
577
580
  embed_onehot,
@@ -604,9 +607,29 @@ class Codebook(Module):
604
607
  ema_inplace(self.cluster_size, cluster_size, self.decay, ema_update_weight)
605
608
  ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
606
609
 
607
- if not self.manual_ema_update:
608
- self.update_ema()
609
- self.expire_codes_(flatten)
610
+ def update_codebook(
611
+ self,
612
+ flatten,
613
+ embed_onehot,
614
+ mask = None,
615
+ ema_update_weight: Tensor | Callable | None = None,
616
+ accum_ema_update = False,
617
+ ema_update = None
618
+ ):
619
+ ema_update = default(ema_update, self.ema_update)
620
+
621
+ if not ema_update and not self.has_dead_code_replacement:
622
+ return
623
+
624
+ self.track_cluster_size_and_embed_avg(flatten, embed_onehot, mask, ema_update_weight, accum_ema_update)
625
+
626
+ if accum_ema_update:
627
+ return
628
+
629
+ if ema_update and not self.manual_ema_update:
630
+ self.update_ema()
631
+
632
+ self.expire_codes_(flatten)
610
633
 
611
634
  def update_ema_indices(
612
635
  self,
@@ -632,7 +655,7 @@ class Codebook(Module):
632
655
  embed_ind = embed_ind.masked_fill(embed_ind == -1, 0)
633
656
  embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
634
657
 
635
- self.update_ema_part(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
658
+ self.update_codebook(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
636
659
 
637
660
  @autocast('cuda', enabled = False)
638
661
  def forward(
@@ -645,7 +668,8 @@ class Codebook(Module):
645
668
  ema_update_weight: Tensor | Callable | None = None,
646
669
  accum_ema_update = False,
647
670
  ema_update = None,
648
- topk = None
671
+ topk = None,
672
+ update_usage = True
649
673
  ):
650
674
  ema_update = default(ema_update, self.ema_update)
651
675
 
@@ -743,8 +767,8 @@ class Codebook(Module):
743
767
  repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
744
768
  quantize = repeated_embed.gather(-2, repeated_embed_ind)
745
769
 
746
- if self.training and ema_update and not freeze_codebook and not exists(topk):
747
- self.update_ema_part(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
770
+ if self.training and update_usage and not freeze_codebook and not exists(topk):
771
+ self.update_codebook(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update, ema_update = ema_update)
748
772
 
749
773
  if needs_codebook_dim:
750
774
  quantize, embed_ind = map(lambda t: rearrange(t, '1 ... -> ...'), (quantize, embed_ind))
@@ -1163,6 +1187,7 @@ class VectorQuantize(Module):
1163
1187
 
1164
1188
  # quantize again
1165
1189
 
1190
+ codebook_forward_kwargs.update(update_usage = False)
1166
1191
  quantize, embed_ind, distances = self._codebook(x, **codebook_forward_kwargs)
1167
1192
 
1168
1193
  if self.training: