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.
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/PKG-INFO +2 -2
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/pyproject.toml +2 -2
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/tests/test_readme.py +4 -2
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/finite_scalar_quantization.py +20 -7
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_fsq.py +22 -11
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/build.yml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/python-publish.yml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/test.yml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.gitignore +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/LICENSE +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/README.md +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_fsq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_lfq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/fsq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/lfq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/simvq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/images/vq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/ruff.toml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/tests/test_latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/tests/test_lfq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/residual_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/sim_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/utils.py +0 -0
- {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.
|
|
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.
|
|
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.
|
|
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.
|
|
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
|
-
|
|
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)
|
|
129
|
-
bounded_z = (z + shift)
|
|
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
|
|
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:
|
|
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 =
|
|
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((
|
|
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]
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/build.yml
RENAMED
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/.github/workflows/test.yml
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_fsq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_lfq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/examples/autoencoder_sim_vq.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/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
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.26.0}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|
|
File without changes
|