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.
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/PKG-INFO +16 -3
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/README.md +15 -2
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/pyproject.toml +1 -1
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/vector_quantize_pytorch.py +35 -10
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/.gitignore +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/LICENSE +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/residual_vq.py +0 -0
- {vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/sim_vq.py +0 -0
- {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.
|
|
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
|
```
|
|
@@ -314,7 +314,7 @@ def rotate_to(src, tgt):
|
|
|
314
314
|
|
|
315
315
|
return inverse(rotated)
|
|
316
316
|
|
|
317
|
-
# directional
|
|
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.
|
|
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
|
|
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
|
-
|
|
608
|
-
|
|
609
|
-
|
|
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.
|
|
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
|
|
747
|
-
self.
|
|
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:
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.28.2 → vector_quantize_pytorch-1.28.3}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|