vector-quantize-pytorch 1.25.2__tar.gz → 1.26.0__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 (33) hide show
  1. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/PKG-INFO +2 -2
  2. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/pyproject.toml +2 -2
  3. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/tests/test_readme.py +4 -2
  4. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/finite_scalar_quantization.py +20 -7
  5. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_fsq.py +22 -11
  6. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/build.yml +0 -0
  7. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/python-publish.yml +0 -0
  8. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/test.yml +0 -0
  9. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.gitignore +0 -0
  10. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/LICENSE +0 -0
  11. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/README.md +0 -0
  12. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder.py +0 -0
  13. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_fsq.py +0 -0
  14. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_lfq.py +0 -0
  15. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_sim_vq.py +0 -0
  16. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/fsq.png +0 -0
  17. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/lfq.png +0 -0
  18. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/simvq.png +0 -0
  19. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/vq.png +0 -0
  20. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/ruff.toml +0 -0
  21. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/tests/test_latent_quantization.py +0 -0
  22. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/tests/test_lfq.py +0 -0
  23. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/__init__.py +0 -0
  24. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/binary_mapper.py +0 -0
  25. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/latent_quantization.py +0 -0
  26. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  27. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  28. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_lfq.py +0 -0
  29. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  30. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_vq.py +0 -0
  31. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/sim_vq.py +0 -0
  32. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/utils.py +0 -0
  33. {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/vector_quantize_pytorch.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.25.2
3
+ Version: 1.26.0
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
@@ -36,7 +36,7 @@ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
36
36
  Requires-Python: >=3.9
37
37
  Requires-Dist: einops>=0.8.0
38
38
  Requires-Dist: einx>=0.3.0
39
- Requires-Dist: torch>=2.0
39
+ Requires-Dist: torch>=2.4
40
40
  Provides-Extra: examples
41
41
  Requires-Dist: torchvision; extra == 'examples'
42
42
  Requires-Dist: tqdm; extra == 'examples'
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.25.2"
3
+ version = "1.26.0"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -23,7 +23,7 @@ classifiers=[
23
23
  ]
24
24
 
25
25
  dependencies = [
26
- "torch>=2.0",
26
+ "torch>=2.4",
27
27
  "einops>=0.8.0",
28
28
  "einx>=0.3.0",
29
29
  ]
@@ -247,13 +247,15 @@ def test_directional_reparam():
247
247
  quantized, indices, _ = rq(x)
248
248
 
249
249
  @pytest.mark.parametrize('preserve_symmetry', (True, False))
250
+ @pytest.mark.parametrize('bound_hard_clamp', (True, False))
250
251
  def test_fsq(
251
- preserve_symmetry
252
+ preserve_symmetry,
253
+ bound_hard_clamp
252
254
  ):
253
255
  from vector_quantize_pytorch import FSQ
254
256
 
255
257
  levels = [8,5,5,5] # see 4.1 and A.4.1 in the paper
256
- quantizer = FSQ(levels, preserve_symmetry = preserve_symmetry)
258
+ quantizer = FSQ(levels, preserve_symmetry = preserve_symmetry, bound_hard_clamp = bound_hard_clamp)
257
259
 
258
260
  x = torch.randn(1, 1024, 4) # 4 since there are 4 levels
259
261
  xhat, indices = quantizer(x)
@@ -11,7 +11,7 @@ from typing import List, Tuple
11
11
  import torch
12
12
  import torch.nn as nn
13
13
  from torch.nn import Module
14
- from torch import tensor, Tensor, int32
14
+ from torch import tensor, Tensor, int32, tanh, atanh, clamp
15
15
  from torch.amp import autocast
16
16
 
17
17
  import einx
@@ -30,6 +30,9 @@ def default(*args):
30
30
  return arg
31
31
  return None
32
32
 
33
+ def identity(t):
34
+ return t
35
+
33
36
  def maybe(fn):
34
37
  @wraps(fn)
35
38
  def inner(x, *args, **kwargs):
@@ -73,6 +76,7 @@ class FSQ(Module):
73
76
  force_quantization_f32 = True,
74
77
  preserve_symmetry = False,
75
78
  noise_dropout = 0.,
79
+ bound_hard_clamp = False # for residual fsq, if input is pre-softclamped to the right range
76
80
  ):
77
81
  super().__init__()
78
82
 
@@ -121,22 +125,31 @@ class FSQ(Module):
121
125
  self.allowed_dtypes = allowed_dtypes
122
126
  self.force_quantization_f32 = force_quantization_f32
123
127
 
124
- def bound(self, z, eps = 1e-3):
128
+ # allow for a hard clamp
129
+
130
+ self.bound_hard_clamp = bound_hard_clamp
131
+
132
+ def bound(self, z, eps = 1e-3, hard_clamp = False):
125
133
  """ Bound `z`, an array of shape (..., d). """
134
+ maybe_tanh = tanh if not hard_clamp else partial(clamp, min = -1., max = 1.)
135
+ maybe_atanh = atanh if not hard_clamp else identity
136
+
126
137
  half_l = (self._levels - 1) * (1 + eps) / 2
127
138
  offset = torch.where(self._levels % 2 == 0, 0.5, 0.0)
128
- shift = (offset / half_l).atanh()
129
- bounded_z = (z + shift).tanh() * half_l - offset
139
+ shift = maybe_atanh(offset / half_l)
140
+ bounded_z = maybe_tanh(z + shift) * half_l - offset
130
141
  half_width = self._levels // 2
131
142
  return round_ste(bounded_z) / half_width
132
143
 
133
144
  # symmetry-preserving and noise-approximated quantization, section 3.2 in https://arxiv.org/abs/2411.19842
134
145
 
135
- def symmetry_preserving_bound(self, z):
146
+ def symmetry_preserving_bound(self, z, hard_clamp = False):
136
147
  """ QL(x) = 2 / (L - 1) * [(L - 1) * (tanh(x) + 1) / 2 + 0.5] - 1 """
148
+ maybe_tanh = tanh if not hard_clamp else partial(clamp, min = -1., max = 1.)
149
+
137
150
  levels_minus_1 = (self._levels - 1)
138
151
  scale = 2. / levels_minus_1
139
- bracket = (levels_minus_1 * (z.tanh() + 1) / 2.) + 0.5
152
+ bracket = (levels_minus_1 * (maybe_tanh(z) + 1) / 2.) + 0.5
140
153
  bracket = floor_ste(bracket)
141
154
  return scale * bracket - 1.
142
155
 
@@ -146,7 +159,7 @@ class FSQ(Module):
146
159
  shape, device, noise_dropout, preserve_symmetry = z.shape[0], z.device, self.noise_dropout, self.preserve_symmetry
147
160
  bound_fn = self.symmetry_preserving_bound if preserve_symmetry else self.bound
148
161
 
149
- bounded_z = bound_fn(z)
162
+ bounded_z = bound_fn(z, hard_clamp = self.bound_hard_clamp)
150
163
 
151
164
  # determine where to add a random offset elementwise
152
165
  # if using noise dropout
@@ -1,11 +1,11 @@
1
+ from __future__ import annotations
2
+
1
3
  import random
2
4
  from math import ceil
3
5
  from functools import partial
4
6
 
5
- from typing import List
6
-
7
7
  import torch
8
- from torch import nn
8
+ from torch import nn, tensor
9
9
  from torch.nn import Module, ModuleList
10
10
  import torch.nn.functional as F
11
11
  from torch.amp import autocast
@@ -52,14 +52,15 @@ class ResidualFSQ(Module):
52
52
  def __init__(
53
53
  self,
54
54
  *,
55
- levels: List[int],
55
+ levels: list[int],
56
56
  num_quantizers,
57
57
  dim = None,
58
58
  is_channel_first = False,
59
59
  quantize_dropout = False,
60
60
  quantize_dropout_cutoff_index = 0,
61
61
  quantize_dropout_multiple_of = 1,
62
- soft_clamp_input_value = None,
62
+ soft_clamp_input_value: float | list[float] | Tensor | None = None,
63
+ bound_hard_clamp = True,
63
64
  **kwargs
64
65
  ):
65
66
  super().__init__()
@@ -74,25 +75,24 @@ class ResidualFSQ(Module):
74
75
  self.is_channel_first = is_channel_first
75
76
  self.num_quantizers = num_quantizers
76
77
 
77
- # soft clamping the input value
78
-
79
- self.soft_clamp_input_value = soft_clamp_input_value
80
-
81
78
  # layers
82
79
 
83
80
  self.levels = levels
84
81
  self.layers = nn.ModuleList([])
85
82
 
86
- levels_tensor = torch.Tensor(levels)
83
+ levels_tensor = tensor(levels)
84
+ assert (levels_tensor > 1).all()
87
85
 
88
86
  scales = []
89
87
 
90
88
  for ind in range(num_quantizers):
91
- scales.append((levels_tensor - 1) ** -ind)
89
+ scales.append(levels_tensor.float() ** -ind)
92
90
 
93
91
  fsq = FSQ(
94
92
  levels = levels,
95
93
  dim = codebook_dim,
94
+ preserve_symmetry = True,
95
+ bound_hard_clamp = bound_hard_clamp,
96
96
  **kwargs
97
97
  )
98
98
 
@@ -111,6 +111,17 @@ class ResidualFSQ(Module):
111
111
  self.quantize_dropout_cutoff_index = quantize_dropout_cutoff_index
112
112
  self.quantize_dropout_multiple_of = quantize_dropout_multiple_of # encodec paper proposes structured dropout, believe this was set to 4
113
113
 
114
+ # soft clamping the input value
115
+
116
+ if bound_hard_clamp:
117
+ assert not exists(soft_clamp_input_value)
118
+ soft_clamp_input_value = 1 + (1 / (levels_tensor - 1))
119
+
120
+ if isinstance(soft_clamp_input_value, (list, float)):
121
+ soft_clamp_input_value = tensor(soft_clamp_input_value)
122
+
123
+ self.register_buffer('soft_clamp_input_value', soft_clamp_input_value, persistent = False)
124
+
114
125
  @property
115
126
  def codebooks(self):
116
127
  codebooks = [layer.implicit_codebook for layer in self.layers]