vector-quantize-pytorch 1.20.2__tar.gz → 1.20.4__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.20.2 → vector_quantize_pytorch-1.20.4}/PKG-INFO +1 -1
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_sim_vq.py +10 -2
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/pyproject.toml +1 -1
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/tests/test_readme.py +14 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/sim_vq.py +23 -5
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/vector_quantize_pytorch.py +9 -9
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/build.yml +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/python-publish.yml +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/test.yml +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.gitignore +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/LICENSE +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/README.md +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_fsq.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_lfq.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/images/fsq.png +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/images/lfq.png +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/images/vq.png +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/ruff.toml +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/tests/test_latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/residual_vq.py +0 -0
- {vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: vector-quantize-pytorch
|
|
3
|
-
Version: 1.20.
|
|
3
|
+
Version: 1.20.4
|
|
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
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_sim_vq.py
RENAMED
|
@@ -14,7 +14,10 @@ lr = 3e-4
|
|
|
14
14
|
train_iter = 10000
|
|
15
15
|
num_codes = 256
|
|
16
16
|
seed = 1234
|
|
17
|
-
|
|
17
|
+
|
|
18
|
+
rotation_trick = True # rotation trick instead ot straight-through
|
|
19
|
+
use_mlp = True # use a one layer mlp with relu instead of linear
|
|
20
|
+
|
|
18
21
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
19
22
|
|
|
20
23
|
def SimVQAutoEncoder(**vq_kwargs):
|
|
@@ -77,7 +80,12 @@ torch.random.manual_seed(seed)
|
|
|
77
80
|
|
|
78
81
|
model = SimVQAutoEncoder(
|
|
79
82
|
codebook_size = num_codes,
|
|
80
|
-
rotation_trick = rotation_trick
|
|
83
|
+
rotation_trick = rotation_trick,
|
|
84
|
+
codebook_transform = nn.Sequential(
|
|
85
|
+
nn.Linear(32, 128),
|
|
86
|
+
nn.ReLU(),
|
|
87
|
+
nn.Linear(128, 32),
|
|
88
|
+
) if use_mlp else None
|
|
81
89
|
).to(device)
|
|
82
90
|
|
|
83
91
|
opt = torch.optim.AdamW(model.parameters(), lr=lr)
|
|
@@ -362,3 +362,17 @@ def test_latent_q():
|
|
|
362
362
|
|
|
363
363
|
assert image_feats.shape == quantized.shape
|
|
364
364
|
assert (quantized == quantizer.indices_to_codes(indices)).all()
|
|
365
|
+
|
|
366
|
+
def test_sim_vq():
|
|
367
|
+
from vector_quantize_pytorch import SimVQ
|
|
368
|
+
|
|
369
|
+
sim_vq = SimVQ(
|
|
370
|
+
dim = 512,
|
|
371
|
+
codebook_size = 1024,
|
|
372
|
+
)
|
|
373
|
+
|
|
374
|
+
x = torch.randn(1, 1024, 512)
|
|
375
|
+
quantized, indices, commit_loss = sim_vq(x)
|
|
376
|
+
|
|
377
|
+
assert x.shape == quantized.shape
|
|
378
|
+
assert torch.allclose(quantized, sim_vq.indices_to_codes(indices), atol = 1e-6)
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
@@ -9,7 +9,7 @@ import torch.nn.functional as F
|
|
|
9
9
|
from einx import get_at
|
|
10
10
|
from einops import rearrange, pack, unpack
|
|
11
11
|
|
|
12
|
-
from vector_quantize_pytorch.vector_quantize_pytorch import
|
|
12
|
+
from vector_quantize_pytorch.vector_quantize_pytorch import rotate_to
|
|
13
13
|
|
|
14
14
|
# helper functions
|
|
15
15
|
|
|
@@ -56,9 +56,9 @@ class SimVQ(Module):
|
|
|
56
56
|
if not exists(codebook_transform):
|
|
57
57
|
codebook_transform = nn.Linear(dim, dim, bias = False)
|
|
58
58
|
|
|
59
|
-
self.
|
|
59
|
+
self.code_transform = codebook_transform
|
|
60
60
|
|
|
61
|
-
self.register_buffer('
|
|
61
|
+
self.register_buffer('frozen_codebook', codebook)
|
|
62
62
|
|
|
63
63
|
|
|
64
64
|
# whether to use rotation trick from Fifty et al.
|
|
@@ -70,6 +70,24 @@ class SimVQ(Module):
|
|
|
70
70
|
|
|
71
71
|
self.input_to_quantize_commit_loss_weight = input_to_quantize_commit_loss_weight
|
|
72
72
|
|
|
73
|
+
@property
|
|
74
|
+
def codebook(self):
|
|
75
|
+
return self.code_transform(self.frozen_codebook)
|
|
76
|
+
|
|
77
|
+
def indices_to_codes(
|
|
78
|
+
self,
|
|
79
|
+
indices
|
|
80
|
+
):
|
|
81
|
+
implicit_codebook = self.codebook
|
|
82
|
+
|
|
83
|
+
frozen_codes = get_at('[c] d, b ... -> b ... d', self.frozen_codebook, indices)
|
|
84
|
+
quantized = self.code_transform(frozen_codes)
|
|
85
|
+
|
|
86
|
+
if self.accept_image_fmap:
|
|
87
|
+
quantized = rearrange(quantized, 'b ... d -> b d ...')
|
|
88
|
+
|
|
89
|
+
return quantized
|
|
90
|
+
|
|
73
91
|
def forward(
|
|
74
92
|
self,
|
|
75
93
|
x
|
|
@@ -78,7 +96,7 @@ class SimVQ(Module):
|
|
|
78
96
|
x = rearrange(x, 'b d h w -> b h w d')
|
|
79
97
|
x, inverse_pack = pack_one(x, 'b * d')
|
|
80
98
|
|
|
81
|
-
implicit_codebook = self.
|
|
99
|
+
implicit_codebook = self.codebook
|
|
82
100
|
|
|
83
101
|
with torch.no_grad():
|
|
84
102
|
dist = torch.cdist(x, implicit_codebook)
|
|
@@ -94,7 +112,7 @@ class SimVQ(Module):
|
|
|
94
112
|
|
|
95
113
|
if self.rotation_trick:
|
|
96
114
|
# rotation trick from @cfifty
|
|
97
|
-
quantized =
|
|
115
|
+
quantized = rotate_to(x, quantized)
|
|
98
116
|
else:
|
|
99
117
|
|
|
100
118
|
commit_loss = (
|
|
@@ -250,21 +250,21 @@ def efficient_rotation_trick_transform(u, q, e):
|
|
|
250
250
|
2 * (e @ rearrange(u, 'b d -> b d 1').detach() @ rearrange(q, 'b d -> b 1 d').detach())
|
|
251
251
|
)
|
|
252
252
|
|
|
253
|
-
def
|
|
253
|
+
def rotate_to(src, tgt):
|
|
254
254
|
# rotation trick STE (https://arxiv.org/abs/2410.06424) to get gradients through VQ layer.
|
|
255
|
-
|
|
256
|
-
|
|
255
|
+
src, inverse = pack_one(src, '* d')
|
|
256
|
+
tgt, _ = pack_one(tgt, '* d')
|
|
257
257
|
|
|
258
|
-
norm_tgt = tgt.norm(dim = -1, keepdim = True)
|
|
259
258
|
norm_src = src.norm(dim = -1, keepdim = True)
|
|
259
|
+
norm_tgt = tgt.norm(dim = -1, keepdim = True)
|
|
260
260
|
|
|
261
|
-
|
|
262
|
-
safe_div(tgt, norm_tgt),
|
|
261
|
+
rotated_tgt = efficient_rotation_trick_transform(
|
|
263
262
|
safe_div(src, norm_src),
|
|
264
|
-
tgt
|
|
263
|
+
safe_div(tgt, norm_tgt),
|
|
264
|
+
src
|
|
265
265
|
).squeeze()
|
|
266
266
|
|
|
267
|
-
rotated =
|
|
267
|
+
rotated = rotated_tgt * safe_div(norm_tgt, norm_src).detach()
|
|
268
268
|
|
|
269
269
|
return inverse(rotated)
|
|
270
270
|
|
|
@@ -1118,7 +1118,7 @@ class VectorQuantize(Module):
|
|
|
1118
1118
|
commit_quantize = maybe_detach(quantize)
|
|
1119
1119
|
|
|
1120
1120
|
if self.rotation_trick:
|
|
1121
|
-
quantize =
|
|
1121
|
+
quantize = rotate_to(x, quantize)
|
|
1122
1122
|
else:
|
|
1123
1123
|
# standard STE to get gradients through VQ layer.
|
|
1124
1124
|
quantize = x + (quantize - x).detach()
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/build.yml
RENAMED
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/.github/workflows/test.yml
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_fsq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/examples/autoencoder_lfq.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/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
|
|
File without changes
|
{vector_quantize_pytorch-1.20.2 → vector_quantize_pytorch-1.20.4}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|