vector-quantize-pytorch 1.20.0__tar.gz → 1.20.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.
Files changed (29) hide show
  1. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/sim_vq.py +21 -11
  4. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/.github/workflows/build.yml +0 -0
  5. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/.github/workflows/python-publish.yml +0 -0
  6. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/.github/workflows/test.yml +0 -0
  7. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/.gitignore +0 -0
  8. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/LICENSE +0 -0
  9. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/README.md +0 -0
  10. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/examples/autoencoder.py +0 -0
  11. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/examples/autoencoder_fsq.py +0 -0
  12. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/examples/autoencoder_lfq.py +0 -0
  13. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/examples/autoencoder_sim_vq.py +0 -0
  14. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/images/fsq.png +0 -0
  15. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/images/lfq.png +0 -0
  16. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/images/vq.png +0 -0
  17. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/ruff.toml +0 -0
  18. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/tests/test_latent_quantization.py +0 -0
  19. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/tests/test_readme.py +0 -0
  20. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/__init__.py +0 -0
  21. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  22. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
  23. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  24. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  25. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
  26. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
  27. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/residual_vq.py +0 -0
  28. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/vector_quantize_pytorch/utils.py +0 -0
  29. {vector_quantize_pytorch-1.20.0 → vector_quantize_pytorch-1.20.2}/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.20.0
3
+ Version: 1.20.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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.20.0"
3
+ version = "1.20.2"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -1,3 +1,4 @@
1
+ from __future__ import annotations
1
2
  from typing import Callable
2
3
 
3
4
  import torch
@@ -38,10 +39,11 @@ class SimVQ(Module):
38
39
  self,
39
40
  dim,
40
41
  codebook_size,
42
+ codebook_transform: Module | None = None,
41
43
  init_fn: Callable = identity,
42
44
  accept_image_fmap = False,
43
- rotation_trick = True, # works even better with rotation trick turned on, with no asymmetric commit loss or straight through
44
- commit_loss_input_to_quantize_weight = 0.25,
45
+ rotation_trick = True, # works even better with rotation trick turned on, with no straight through and the commit loss from input to quantize
46
+ input_to_quantize_commit_loss_weight = 0.25,
45
47
  ):
46
48
  super().__init__()
47
49
  self.accept_image_fmap = accept_image_fmap
@@ -51,7 +53,11 @@ class SimVQ(Module):
51
53
 
52
54
  # the codebook is actually implicit from a linear layer from frozen gaussian or uniform
53
55
 
54
- self.codebook_to_codes = nn.Linear(dim, dim, bias = False)
56
+ if not exists(codebook_transform):
57
+ codebook_transform = nn.Linear(dim, dim, bias = False)
58
+
59
+ self.codebook_to_codes = codebook_transform
60
+
55
61
  self.register_buffer('codebook', codebook)
56
62
 
57
63
 
@@ -59,11 +65,10 @@ class SimVQ(Module):
59
65
  # https://arxiv.org/abs/2410.06424
60
66
 
61
67
  self.rotation_trick = rotation_trick
62
- self.register_buffer('zero', torch.tensor(0.), persistent = False)
63
68
 
64
69
  # commit loss weighting - weighing input to quantize a bit less is crucial for it to work
65
70
 
66
- self.commit_loss_input_to_quantize_weight = commit_loss_input_to_quantize_weight
71
+ self.input_to_quantize_commit_loss_weight = input_to_quantize_commit_loss_weight
67
72
 
68
73
  def forward(
69
74
  self,
@@ -83,18 +88,18 @@ class SimVQ(Module):
83
88
 
84
89
  quantized = get_at('[c] d, b n -> b n d', implicit_codebook, indices)
85
90
 
91
+ # commit loss and straight through, as was done in the paper
92
+
93
+ commit_loss = F.mse_loss(x.detach(), quantized)
94
+
86
95
  if self.rotation_trick:
87
96
  # rotation trick from @cfifty
88
-
89
97
  quantized = rotate_from_to(quantized, x)
90
-
91
- commit_loss = self.zero
92
98
  else:
93
- # commit loss and straight through, as was done in the paper
94
99
 
95
100
  commit_loss = (
96
- F.mse_loss(x, quantized.detach()) * self.commit_loss_input_to_quantize_weight +
97
- F.mse_loss(x.detach(), quantized)
101
+ commit_loss +
102
+ F.mse_loss(x, quantized.detach()) * self.input_to_quantize_commit_loss_weight
98
103
  )
99
104
 
100
105
  quantized = (quantized - x).detach() + x
@@ -115,6 +120,11 @@ if __name__ == '__main__':
115
120
 
116
121
  sim_vq = SimVQ(
117
122
  dim = 512,
123
+ codebook_transform = nn.Sequential(
124
+ nn.Linear(512, 1024),
125
+ nn.ReLU(),
126
+ nn.Linear(1024, 512)
127
+ ),
118
128
  codebook_size = 1024,
119
129
  accept_image_fmap = True
120
130
  )