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.
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/PKG-INFO +13 -1
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/README.md +12 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/pyproject.toml +1 -1
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/finite_scalar_quantization.py +22 -1
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/lookup_free_quantization.py +22 -1
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +2 -1
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/.gitignore +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/LICENSE +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/residual_vq.py +0 -0
- {vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/sim_vq.py +0 -0
- {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.
|
|
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
|
+
```
|
|
@@ -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
|
|
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
|
|
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
|
|
|
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.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.28.0 → vector_quantize_pytorch-1.28.2}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|