vector-quantize-pytorch 1.18.0__tar.gz → 1.18.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.18.0 → vector_quantize_pytorch-1.18.2}/PKG-INFO +1 -1
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/pyproject.toml +1 -1
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/tests/test_readme.py +5 -2
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/lookup_free_quantization.py +1 -1
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +37 -26
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/build.yml +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/python-publish.yml +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/test.yml +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.gitignore +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/LICENSE +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/README.md +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder_fsq.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder_lfq.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/images/fsq.png +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/images/lfq.png +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/images/vq.png +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/ruff.toml +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/tests/test_latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/vector_quantize_pytorch/residual_vq.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: vector-quantize-pytorch
|
|
3
|
-
Version: 1.18.
|
|
3
|
+
Version: 1.18.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
|
|
@@ -5,8 +5,10 @@ def exists(v):
|
|
|
5
5
|
return v is not None
|
|
6
6
|
|
|
7
7
|
@pytest.mark.parametrize('use_cosine_sim', (True, False))
|
|
8
|
+
@pytest.mark.parametrize('rotation_trick', (True, False))
|
|
8
9
|
def test_vq(
|
|
9
|
-
use_cosine_sim
|
|
10
|
+
use_cosine_sim,
|
|
11
|
+
rotation_trick
|
|
10
12
|
):
|
|
11
13
|
from vector_quantize_pytorch import VectorQuantize
|
|
12
14
|
|
|
@@ -15,7 +17,8 @@ def test_vq(
|
|
|
15
17
|
codebook_size = 512, # codebook size
|
|
16
18
|
decay = 0.8, # the exponential moving average decay, lower means the dictionary will change faster
|
|
17
19
|
commitment_weight = 1., # the weight on the commitment loss
|
|
18
|
-
use_cosine_sim = use_cosine_sim
|
|
20
|
+
use_cosine_sim = use_cosine_sim,
|
|
21
|
+
rotation_trick = rotation_trick
|
|
19
22
|
)
|
|
20
23
|
|
|
21
24
|
x = torch.randn(1, 1024, 256)
|
|
@@ -100,7 +100,7 @@ class LFQ(Module):
|
|
|
100
100
|
dim = None,
|
|
101
101
|
codebook_size = None,
|
|
102
102
|
entropy_loss_weight = 0.1,
|
|
103
|
-
commitment_loss_weight = 0
|
|
103
|
+
commitment_loss_weight = 0.,
|
|
104
104
|
diversity_gamma = 1.,
|
|
105
105
|
straight_through_activation = nn.Identity(),
|
|
106
106
|
num_codebooks = 1,
|
|
@@ -28,8 +28,11 @@ def noop(*args, **kwargs):
|
|
|
28
28
|
def identity(t):
|
|
29
29
|
return t
|
|
30
30
|
|
|
31
|
-
def l2norm(t):
|
|
32
|
-
return F.normalize(t, p = 2, dim =
|
|
31
|
+
def l2norm(t, dim = -1, eps = 1e-6):
|
|
32
|
+
return F.normalize(t, p = 2, dim = dim, eps = eps)
|
|
33
|
+
|
|
34
|
+
def safe_div(num, den, eps = 1e-6):
|
|
35
|
+
return num / den.clamp(min = eps)
|
|
33
36
|
|
|
34
37
|
def Sequential(*modules):
|
|
35
38
|
modules = [*filter(exists, modules)]
|
|
@@ -73,6 +76,19 @@ def lens_to_mask(lens, max_length):
|
|
|
73
76
|
seq = torch.arange(max_length, device = lens.device)
|
|
74
77
|
return seq < lens[:, None]
|
|
75
78
|
|
|
79
|
+
def efficient_rotation_trick_transform(u, q, e):
|
|
80
|
+
"""
|
|
81
|
+
4.2 in https://arxiv.org/abs/2410.06424
|
|
82
|
+
"""
|
|
83
|
+
e = rearrange(e, 'b d -> b 1 d')
|
|
84
|
+
w = l2norm(u + q, dim = 1).detach()
|
|
85
|
+
|
|
86
|
+
return (
|
|
87
|
+
e -
|
|
88
|
+
2 * (e @ rearrange(w, 'b d -> b d 1') @ rearrange(w, 'b d -> b 1 d')) +
|
|
89
|
+
2 * (e @ rearrange(u, 'b d -> b d 1').detach() @ rearrange(q, 'b d -> b 1 d').detach())
|
|
90
|
+
)
|
|
91
|
+
|
|
76
92
|
def uniform_init(*shape):
|
|
77
93
|
t = torch.empty(shape)
|
|
78
94
|
nn.init.kaiming_uniform_(t)
|
|
@@ -811,7 +827,7 @@ class VectorQuantize(Module):
|
|
|
811
827
|
stochastic_sample_codes = False,
|
|
812
828
|
sample_codebook_temp = 1.,
|
|
813
829
|
straight_through = False,
|
|
814
|
-
rotation_trick = True, # Propagate grads through VQ layer w/ rotation trick: https://arxiv.org/abs/2410.06424
|
|
830
|
+
rotation_trick = True, # Propagate grads through VQ layer w/ rotation trick: https://arxiv.org/abs/2410.06424 by @cfifty
|
|
815
831
|
reinmax = False, # using reinmax for improved straight-through, assuming straight through helps at all
|
|
816
832
|
sync_codebook = None,
|
|
817
833
|
sync_affine_param = False,
|
|
@@ -946,13 +962,6 @@ class VectorQuantize(Module):
|
|
|
946
962
|
|
|
947
963
|
self._codebook.embed.copy_(codes)
|
|
948
964
|
|
|
949
|
-
@staticmethod
|
|
950
|
-
def rotation_trick_transform(u, q, e):
|
|
951
|
-
w = ((u + q) / torch.norm(u + q, dim=1, keepdim=True)).detach()
|
|
952
|
-
e = e - 2 * torch.bmm(torch.bmm(e, w.unsqueeze(-1)), w.unsqueeze(1)) + 2 * torch.bmm(
|
|
953
|
-
torch.bmm(e, u.unsqueeze(-1).detach()), q.unsqueeze(1).detach())
|
|
954
|
-
return e
|
|
955
|
-
|
|
956
965
|
def get_codes_from_indices(self, indices):
|
|
957
966
|
codebook = self.codebook
|
|
958
967
|
is_multiheaded = codebook.ndim > 2
|
|
@@ -1103,23 +1112,25 @@ class VectorQuantize(Module):
|
|
|
1103
1112
|
|
|
1104
1113
|
commit_quantize = maybe_detach(quantize)
|
|
1105
1114
|
|
|
1106
|
-
# Use the rotation trick (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
|
|
1107
1115
|
if self.rotation_trick:
|
|
1108
|
-
|
|
1109
|
-
x = x
|
|
1110
|
-
quantize = quantize
|
|
1111
|
-
|
|
1112
|
-
|
|
1113
|
-
|
|
1114
|
-
|
|
1115
|
-
|
|
1116
|
-
x
|
|
1117
|
-
|
|
1118
|
-
|
|
1119
|
-
|
|
1120
|
-
|
|
1121
|
-
quantize =
|
|
1122
|
-
|
|
1116
|
+
# rotation trick STE (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
|
|
1117
|
+
x, inverse = pack_one(x, '* d')
|
|
1118
|
+
quantize, _ = pack_one(quantize, '* d')
|
|
1119
|
+
|
|
1120
|
+
norm_x = x.norm(dim = -1, keepdim = True)
|
|
1121
|
+
norm_quantize = quantize.norm(dim = -1, keepdim = True)
|
|
1122
|
+
|
|
1123
|
+
rot_quantize = efficient_rotation_trick_transform(
|
|
1124
|
+
safe_div(x, norm_x),
|
|
1125
|
+
safe_div(quantize, norm_quantize),
|
|
1126
|
+
x
|
|
1127
|
+
).squeeze()
|
|
1128
|
+
|
|
1129
|
+
quantize = rot_quantize * safe_div(norm_quantize, norm_x).detach()
|
|
1130
|
+
|
|
1131
|
+
x, quantize = inverse(x), inverse(quantize)
|
|
1132
|
+
else:
|
|
1133
|
+
# standard STE to get gradients through VQ layer.
|
|
1123
1134
|
quantize = x + (quantize - x).detach()
|
|
1124
1135
|
|
|
1125
1136
|
if self.sync_update_v > 0.:
|
{vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/build.yml
RENAMED
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/.github/workflows/test.yml
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder_fsq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/examples/autoencoder_lfq.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.18.0 → vector_quantize_pytorch-1.18.2}/tests/test_latent_quantization.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|