vector-quantize-pytorch 1.19.3__tar.gz → 1.19.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.19.3 → vector_quantize_pytorch-1.19.4}/PKG-INFO +1 -1
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder.py +20 -35
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder_fsq.py +19 -32
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder_lfq.py +27 -42
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/pyproject.toml +1 -1
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/__init__.py +2 -0
- vector_quantize_pytorch-1.19.4/vector_quantize_pytorch/utils.py +57 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/build.yml +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/python-publish.yml +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/test.yml +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.gitignore +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/LICENSE +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/README.md +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/images/fsq.png +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/images/lfq.png +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/images/vq.png +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/ruff.toml +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/tests/test_latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/tests/test_readme.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/residual_vq.py +0 -0
- {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/vector_quantize_pytorch.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: vector-quantize-pytorch
|
|
3
|
-
Version: 1.19.
|
|
3
|
+
Version: 1.19.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
|
|
@@ -8,45 +8,29 @@ import torch.nn as nn
|
|
|
8
8
|
from torchvision import datasets, transforms
|
|
9
9
|
from torch.utils.data import DataLoader
|
|
10
10
|
|
|
11
|
-
from vector_quantize_pytorch import VectorQuantize
|
|
12
|
-
|
|
11
|
+
from vector_quantize_pytorch import VectorQuantize, Sequential
|
|
13
12
|
|
|
14
13
|
lr = 3e-4
|
|
15
14
|
train_iter = 1000
|
|
16
15
|
num_codes = 256
|
|
17
16
|
seed = 1234
|
|
17
|
+
rotation_trick = True
|
|
18
18
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
19
19
|
|
|
20
|
-
|
|
21
|
-
|
|
22
|
-
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
|
|
26
|
-
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
nn.GELU(),
|
|
35
|
-
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
36
|
-
nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),
|
|
37
|
-
]
|
|
38
|
-
)
|
|
39
|
-
return
|
|
40
|
-
|
|
41
|
-
def forward(self, x):
|
|
42
|
-
for layer in self.layers:
|
|
43
|
-
if isinstance(layer, VectorQuantize):
|
|
44
|
-
x, indices, commit_loss = layer(x)
|
|
45
|
-
else:
|
|
46
|
-
x = layer(x)
|
|
47
|
-
|
|
48
|
-
return x.clamp(-1, 1), indices, commit_loss
|
|
49
|
-
|
|
20
|
+
def SimpleVQAutoEncoder(**vq_kwargs):
|
|
21
|
+
return Sequential(
|
|
22
|
+
nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1),
|
|
23
|
+
nn.MaxPool2d(kernel_size=2, stride=2),
|
|
24
|
+
nn.GELU(),
|
|
25
|
+
nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),
|
|
26
|
+
nn.MaxPool2d(kernel_size=2, stride=2),
|
|
27
|
+
VectorQuantize(dim=32, accept_image_fmap = True, **vq_kwargs),
|
|
28
|
+
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
29
|
+
nn.Conv2d(32, 16, kernel_size=3, stride=1, padding=1),
|
|
30
|
+
nn.GELU(),
|
|
31
|
+
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
32
|
+
nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),
|
|
33
|
+
)
|
|
50
34
|
|
|
51
35
|
def train(model, train_loader, train_iterations=1000, alpha=10):
|
|
52
36
|
def iterate_dataset(data_loader):
|
|
@@ -62,7 +46,10 @@ def train(model, train_loader, train_iterations=1000, alpha=10):
|
|
|
62
46
|
for _ in (pbar := trange(train_iterations)):
|
|
63
47
|
opt.zero_grad()
|
|
64
48
|
x, _ = next(iterate_dataset(train_loader))
|
|
49
|
+
|
|
65
50
|
out, indices, cmt_loss = model(x)
|
|
51
|
+
out = out.clamp(-1., 1.)
|
|
52
|
+
|
|
66
53
|
rec_loss = (out - x).abs().mean()
|
|
67
54
|
(rec_loss + alpha * cmt_loss).backward()
|
|
68
55
|
|
|
@@ -72,8 +59,6 @@ def train(model, train_loader, train_iterations=1000, alpha=10):
|
|
|
72
59
|
+ f"cmt loss: {cmt_loss.item():.3f} | "
|
|
73
60
|
+ f"active %: {indices.unique().numel() / num_codes * 100:.3f}"
|
|
74
61
|
)
|
|
75
|
-
return
|
|
76
|
-
|
|
77
62
|
|
|
78
63
|
transform = transforms.Compose(
|
|
79
64
|
[transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))]
|
|
@@ -91,7 +76,7 @@ torch.random.manual_seed(seed)
|
|
|
91
76
|
|
|
92
77
|
model = SimpleVQAutoEncoder(
|
|
93
78
|
codebook_size=num_codes,
|
|
94
|
-
rotation_trick=
|
|
79
|
+
rotation_trick=rotation_trick
|
|
95
80
|
).to(device)
|
|
96
81
|
|
|
97
82
|
opt = torch.optim.AdamW(model.parameters(), lr=lr)
|
{vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder_fsq.py
RENAMED
|
@@ -9,7 +9,7 @@ import torch.nn as nn
|
|
|
9
9
|
from torchvision import datasets, transforms
|
|
10
10
|
from torch.utils.data import DataLoader
|
|
11
11
|
|
|
12
|
-
from vector_quantize_pytorch import FSQ
|
|
12
|
+
from vector_quantize_pytorch import FSQ, Sequential
|
|
13
13
|
|
|
14
14
|
|
|
15
15
|
lr = 3e-4
|
|
@@ -20,36 +20,22 @@ seed = 1234
|
|
|
20
20
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
21
21
|
|
|
22
22
|
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
|
|
26
|
-
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
40
|
-
nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),
|
|
41
|
-
]
|
|
42
|
-
)
|
|
43
|
-
return
|
|
44
|
-
|
|
45
|
-
def forward(self, x):
|
|
46
|
-
for layer in self.layers:
|
|
47
|
-
if isinstance(layer, FSQ):
|
|
48
|
-
x, indices = layer(x)
|
|
49
|
-
else:
|
|
50
|
-
x = layer(x)
|
|
51
|
-
|
|
52
|
-
return x.clamp(-1, 1), indices
|
|
23
|
+
def SimpleFSQAutoEncoder(levels: list[int]):
|
|
24
|
+
return Sequential(
|
|
25
|
+
nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1),
|
|
26
|
+
nn.MaxPool2d(kernel_size=2, stride=2),
|
|
27
|
+
nn.GELU(),
|
|
28
|
+
nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),
|
|
29
|
+
nn.MaxPool2d(kernel_size=2, stride=2),
|
|
30
|
+
nn.Conv2d(32, len(levels), kernel_size=1),
|
|
31
|
+
FSQ(levels),
|
|
32
|
+
nn.Conv2d(len(levels), 32, kernel_size=3, stride=1, padding=1),
|
|
33
|
+
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
34
|
+
nn.Conv2d(32, 16, kernel_size=3, stride=1, padding=1),
|
|
35
|
+
nn.GELU(),
|
|
36
|
+
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
37
|
+
nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),
|
|
38
|
+
)
|
|
53
39
|
|
|
54
40
|
|
|
55
41
|
def train(model, train_loader, train_iterations=1000):
|
|
@@ -67,6 +53,8 @@ def train(model, train_loader, train_iterations=1000):
|
|
|
67
53
|
opt.zero_grad()
|
|
68
54
|
x, _ = next(iterate_dataset(train_loader))
|
|
69
55
|
out, indices = model(x)
|
|
56
|
+
out = out.clamp(-1., 1.)
|
|
57
|
+
|
|
70
58
|
rec_loss = (out - x).abs().mean()
|
|
71
59
|
rec_loss.backward()
|
|
72
60
|
|
|
@@ -75,7 +63,6 @@ def train(model, train_loader, train_iterations=1000):
|
|
|
75
63
|
f"rec loss: {rec_loss.item():.3f} | "
|
|
76
64
|
+ f"active %: {indices.unique().numel() / num_codes * 100:.3f}"
|
|
77
65
|
)
|
|
78
|
-
return
|
|
79
66
|
|
|
80
67
|
|
|
81
68
|
transform = transforms.Compose(
|
{vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder_lfq.py
RENAMED
|
@@ -10,7 +10,7 @@ import torch.nn.functional as F
|
|
|
10
10
|
from torchvision import datasets, transforms
|
|
11
11
|
from torch.utils.data import DataLoader
|
|
12
12
|
|
|
13
|
-
from vector_quantize_pytorch import LFQ
|
|
13
|
+
from vector_quantize_pytorch import LFQ, Sequential
|
|
14
14
|
|
|
15
15
|
lr = 3e-4
|
|
16
16
|
train_iter = 1000
|
|
@@ -22,46 +22,31 @@ spherical = True
|
|
|
22
22
|
|
|
23
23
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
24
24
|
|
|
25
|
-
|
|
26
|
-
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
)
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
42
|
-
|
|
43
|
-
|
|
44
|
-
|
|
45
|
-
)
|
|
46
|
-
|
|
47
|
-
|
|
48
|
-
|
|
49
|
-
|
|
50
|
-
nn.Conv2d(quantize_dim, 32, kernel_size=3, stride=1, padding=1),
|
|
51
|
-
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
52
|
-
nn.Conv2d(32, 16, kernel_size=3, stride=1, padding=1),
|
|
53
|
-
nn.GELU(),
|
|
54
|
-
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
55
|
-
nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),
|
|
56
|
-
)
|
|
57
|
-
return
|
|
58
|
-
|
|
59
|
-
def forward(self, x):
|
|
60
|
-
x = self.encode(x)
|
|
61
|
-
x, indices, entropy_aux_loss = self.quantize(x)
|
|
62
|
-
x = self.decode(x)
|
|
63
|
-
return x.clamp(-1, 1), indices, entropy_aux_loss
|
|
64
|
-
|
|
25
|
+
def LFQAutoEncoder(
|
|
26
|
+
codebook_size,
|
|
27
|
+
**vq_kwargs
|
|
28
|
+
):
|
|
29
|
+
assert log2(codebook_size).is_integer()
|
|
30
|
+
quantize_dim = int(log2(codebook_size))
|
|
31
|
+
|
|
32
|
+
return Sequential(
|
|
33
|
+
nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1),
|
|
34
|
+
nn.MaxPool2d(kernel_size=2, stride=2),
|
|
35
|
+
nn.GELU(),
|
|
36
|
+
nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),
|
|
37
|
+
nn.MaxPool2d(kernel_size=2, stride=2),
|
|
38
|
+
# In general norm layers are commonly used in Resnet-based encoder/decoders
|
|
39
|
+
# explicitly add one here with affine=False to avoid introducing new parameters
|
|
40
|
+
nn.GroupNorm(4, 32, affine=False),
|
|
41
|
+
nn.Conv2d(32, quantize_dim, kernel_size=1),
|
|
42
|
+
LFQ(dim=quantize_dim, **vq_kwargs),
|
|
43
|
+
nn.Conv2d(quantize_dim, 32, kernel_size=3, stride=1, padding=1),
|
|
44
|
+
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
45
|
+
nn.Conv2d(32, 16, kernel_size=3, stride=1, padding=1),
|
|
46
|
+
nn.GELU(),
|
|
47
|
+
nn.Upsample(scale_factor=2, mode="nearest"),
|
|
48
|
+
nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),
|
|
49
|
+
)
|
|
65
50
|
|
|
66
51
|
def train(model, train_loader, train_iterations=1000):
|
|
67
52
|
def iterate_dataset(data_loader):
|
|
@@ -78,6 +63,7 @@ def train(model, train_loader, train_iterations=1000):
|
|
|
78
63
|
opt.zero_grad()
|
|
79
64
|
x, _ = next(iterate_dataset(train_loader))
|
|
80
65
|
out, indices, entropy_aux_loss = model(x)
|
|
66
|
+
out = out.clamp(-1., 1.)
|
|
81
67
|
|
|
82
68
|
rec_loss = F.l1_loss(out, x)
|
|
83
69
|
(rec_loss + entropy_aux_loss).backward()
|
|
@@ -88,7 +74,6 @@ def train(model, train_loader, train_iterations=1000):
|
|
|
88
74
|
+ f"entropy aux loss: {entropy_aux_loss.item():.3f} | "
|
|
89
75
|
+ f"active %: {indices.unique().numel() / codebook_size * 100:.3f}"
|
|
90
76
|
)
|
|
91
|
-
return
|
|
92
77
|
|
|
93
78
|
transform = transforms.Compose(
|
|
94
79
|
[transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))]
|
|
@@ -6,3 +6,5 @@ from vector_quantize_pytorch.lookup_free_quantization import LFQ
|
|
|
6
6
|
from vector_quantize_pytorch.residual_lfq import ResidualLFQ, GroupedResidualLFQ
|
|
7
7
|
from vector_quantize_pytorch.residual_fsq import ResidualFSQ, GroupedResidualFSQ
|
|
8
8
|
from vector_quantize_pytorch.latent_quantization import LatentQuantize
|
|
9
|
+
|
|
10
|
+
from vector_quantize_pytorch.utils import Sequential
|
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
from torch import nn
|
|
3
|
+
from torch.nn import Module, ModuleList
|
|
4
|
+
|
|
5
|
+
# quantization
|
|
6
|
+
|
|
7
|
+
from vector_quantize_pytorch.vector_quantize_pytorch import VectorQuantize
|
|
8
|
+
from vector_quantize_pytorch.residual_vq import ResidualVQ, GroupedResidualVQ
|
|
9
|
+
from vector_quantize_pytorch.random_projection_quantizer import RandomProjectionQuantizer
|
|
10
|
+
from vector_quantize_pytorch.finite_scalar_quantization import FSQ
|
|
11
|
+
from vector_quantize_pytorch.lookup_free_quantization import LFQ
|
|
12
|
+
from vector_quantize_pytorch.residual_lfq import ResidualLFQ, GroupedResidualLFQ
|
|
13
|
+
from vector_quantize_pytorch.residual_fsq import ResidualFSQ, GroupedResidualFSQ
|
|
14
|
+
from vector_quantize_pytorch.latent_quantization import LatentQuantize
|
|
15
|
+
|
|
16
|
+
QUANTIZE_KLASSES = (
|
|
17
|
+
VectorQuantize,
|
|
18
|
+
ResidualVQ,
|
|
19
|
+
GroupedResidualVQ,
|
|
20
|
+
RandomProjectionQuantizer,
|
|
21
|
+
FSQ,
|
|
22
|
+
LFQ,
|
|
23
|
+
ResidualLFQ,
|
|
24
|
+
GroupedResidualLFQ,
|
|
25
|
+
ResidualFSQ,
|
|
26
|
+
GroupedResidualFSQ,
|
|
27
|
+
LatentQuantize
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
# classes
|
|
31
|
+
|
|
32
|
+
class Sequential(Module):
|
|
33
|
+
def __init__(
|
|
34
|
+
self,
|
|
35
|
+
*fns: Module
|
|
36
|
+
):
|
|
37
|
+
super().__init__()
|
|
38
|
+
assert sum([int(isinstance(fn, QUANTIZE_KLASSES)) for fn in fns]) == 1, 'this special Sequential must contain exactly one quantizer'
|
|
39
|
+
|
|
40
|
+
self.fns = ModuleList(fns)
|
|
41
|
+
|
|
42
|
+
def forward(
|
|
43
|
+
self,
|
|
44
|
+
x,
|
|
45
|
+
**kwargs
|
|
46
|
+
):
|
|
47
|
+
for fn in self.fns:
|
|
48
|
+
|
|
49
|
+
if not isinstance(fn, QUANTIZE_KLASSES):
|
|
50
|
+
x = fn(x)
|
|
51
|
+
continue
|
|
52
|
+
|
|
53
|
+
x, *rest = fn(x, **kwargs)
|
|
54
|
+
|
|
55
|
+
output = (x, *rest)
|
|
56
|
+
|
|
57
|
+
return output
|
{vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/build.yml
RENAMED
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/test.yml
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
|
{vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.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
|
|
File without changes
|