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.
Files changed (27) hide show
  1. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder.py +20 -35
  3. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder_fsq.py +19 -32
  4. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/examples/autoencoder_lfq.py +27 -42
  5. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/pyproject.toml +1 -1
  6. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/__init__.py +2 -0
  7. vector_quantize_pytorch-1.19.4/vector_quantize_pytorch/utils.py +57 -0
  8. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/build.yml +0 -0
  9. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/python-publish.yml +0 -0
  10. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.github/workflows/test.yml +0 -0
  11. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/.gitignore +0 -0
  12. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/LICENSE +0 -0
  13. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/README.md +0 -0
  14. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/images/fsq.png +0 -0
  15. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/images/lfq.png +0 -0
  16. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/images/vq.png +0 -0
  17. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/ruff.toml +0 -0
  18. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/tests/test_latent_quantization.py +0 -0
  19. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/tests/test_readme.py +0 -0
  20. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  21. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
  22. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  23. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  24. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
  25. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
  26. {vector_quantize_pytorch-1.19.3 → vector_quantize_pytorch-1.19.4}/vector_quantize_pytorch/residual_vq.py +0 -0
  27. {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
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
- class SimpleVQAutoEncoder(nn.Module):
22
- def __init__(self, **vq_kwargs):
23
- super().__init__()
24
- self.layers = nn.ModuleList(
25
- [
26
- nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1),
27
- nn.MaxPool2d(kernel_size=2, stride=2),
28
- nn.GELU(),
29
- nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),
30
- nn.MaxPool2d(kernel_size=2, stride=2),
31
- VectorQuantize(dim=32, accept_image_fmap = True, **vq_kwargs),
32
- nn.Upsample(scale_factor=2, mode="nearest"),
33
- nn.Conv2d(32, 16, kernel_size=3, stride=1, padding=1),
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=True
79
+ rotation_trick=rotation_trick
95
80
  ).to(device)
96
81
 
97
82
  opt = torch.optim.AdamW(model.parameters(), lr=lr)
@@ -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
- class SimpleFSQAutoEncoder(nn.Module):
24
- def __init__(self, levels: list[int]):
25
- super().__init__()
26
- self.layers = nn.ModuleList(
27
- [
28
- nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1),
29
- nn.MaxPool2d(kernel_size=2, stride=2),
30
- nn.GELU(),
31
- nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),
32
- nn.MaxPool2d(kernel_size=2, stride=2),
33
- nn.Conv2d(32, len(levels), kernel_size=1),
34
- FSQ(levels),
35
- nn.Conv2d(len(levels), 32, kernel_size=3, stride=1, padding=1),
36
- nn.Upsample(scale_factor=2, mode="nearest"),
37
- nn.Conv2d(32, 16, kernel_size=3, stride=1, padding=1),
38
- nn.GELU(),
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(
@@ -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
- class LFQAutoEncoder(nn.Module):
26
- def __init__(
27
- self,
28
- codebook_size,
29
- **vq_kwargs
30
- ):
31
- super().__init__()
32
- assert log2(codebook_size).is_integer()
33
- quantize_dim = int(log2(codebook_size))
34
-
35
- self.encode = nn.Sequential(
36
- nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1),
37
- nn.MaxPool2d(kernel_size=2, stride=2),
38
- nn.GELU(),
39
- nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),
40
- nn.MaxPool2d(kernel_size=2, stride=2),
41
- # In general norm layers are commonly used in Resnet-based encoder/decoders
42
- # explicitly add one here with affine=False to avoid introducing new parameters
43
- nn.GroupNorm(4, 32, affine=False),
44
- nn.Conv2d(32, quantize_dim, kernel_size=1),
45
- )
46
-
47
- self.quantize = LFQ(dim=quantize_dim, **vq_kwargs)
48
-
49
- self.decode = nn.Sequential(
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,))]
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.19.3"
3
+ version = "1.19.4"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -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