vector-quantize-pytorch 1.25.2__tar.gz → 1.27.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.27.0}/PKG-INFO +2 -2
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/pyproject.toml +2 -2
- vector_quantize_pytorch-1.27.0/tests/test_beam.py +64 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/tests/test_readme.py +4 -2
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/finite_scalar_quantization.py +20 -7
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_fsq.py +22 -11
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_vq.py +180 -27
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/vector_quantize_pytorch.py +176 -278
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.github/workflows/build.yml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.github/workflows/python-publish.yml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.github/workflows/test.yml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/.gitignore +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/LICENSE +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/README.md +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_fsq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_lfq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/fsq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/lfq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/simvq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/images/vq.png +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/ruff.toml +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/tests/test_latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/tests/test_lfq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/sim_vq.py +0 -0
- {vector_quantize_pytorch-1.25.2 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/utils.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.27.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.27.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
|
]
|
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
from vector_quantize_pytorch import VectorQuantize
|
|
3
|
+
|
|
4
|
+
def test_topk_and_manual_ema_update():
|
|
5
|
+
|
|
6
|
+
vq1 = VectorQuantize(
|
|
7
|
+
dim = 256,
|
|
8
|
+
codebook_size = 512
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
vq2 = VectorQuantize(
|
|
12
|
+
dim = 256,
|
|
13
|
+
codebook_size = 512
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
vq2.load_state_dict(vq1.state_dict())
|
|
17
|
+
|
|
18
|
+
x = torch.randn(1, 1024, 256)
|
|
19
|
+
mask = torch.randint(0, 2, (1, 1024)).bool()
|
|
20
|
+
|
|
21
|
+
vq1.train()
|
|
22
|
+
quantize1, indices1, commit_loss1 = vq1(x, mask = mask)
|
|
23
|
+
|
|
24
|
+
vq2.train()
|
|
25
|
+
quantize2, indices2, commit_losses = vq2(x, mask = mask, topk = 1, ema_update = False)
|
|
26
|
+
|
|
27
|
+
assert quantize2.shape == (1, 1024, 1, 256)
|
|
28
|
+
assert indices2.shape == (1, 1024, 1)
|
|
29
|
+
assert commit_losses.shape == (1, 1024, 1)
|
|
30
|
+
|
|
31
|
+
top_quantize2 = quantize2[..., 0, :]
|
|
32
|
+
top_indices2 = indices2[..., 0]
|
|
33
|
+
|
|
34
|
+
assert torch.allclose(commit_loss1, commit_losses.sum() / mask.sum())
|
|
35
|
+
assert torch.equal(indices1, top_indices2)
|
|
36
|
+
assert torch.allclose(quantize1, top_quantize2)
|
|
37
|
+
|
|
38
|
+
assert not torch.allclose(vq1._codebook.embed_avg, vq2._codebook.embed_avg)
|
|
39
|
+
|
|
40
|
+
vq2.update_ema_indices(x, top_indices2, mask = mask)
|
|
41
|
+
|
|
42
|
+
assert torch.allclose(vq1._codebook.cluster_size, vq2._codebook.cluster_size)
|
|
43
|
+
assert torch.allclose(vq1._codebook.embed_avg, vq2._codebook.embed_avg)
|
|
44
|
+
assert torch.allclose(vq1.codebook, vq2.codebook)
|
|
45
|
+
|
|
46
|
+
def test_beam_search():
|
|
47
|
+
import torch
|
|
48
|
+
from vector_quantize_pytorch import ResidualVQ
|
|
49
|
+
|
|
50
|
+
residual_vq = ResidualVQ(
|
|
51
|
+
dim = 256,
|
|
52
|
+
num_quantizers = 8, # specify number of quantizers
|
|
53
|
+
codebook_size = 1024, # codebook size
|
|
54
|
+
quantize_dropout = True
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
x = torch.randn(1, 1024, 256)
|
|
58
|
+
|
|
59
|
+
for _ in range(5):
|
|
60
|
+
quantized, indices, commit_loss = residual_vq(x, beam_size = 3)
|
|
61
|
+
|
|
62
|
+
assert quantized.shape == (1, 1024, 256)
|
|
63
|
+
assert indices.shape == (1, 1024, 8)
|
|
64
|
+
assert commit_loss.shape == (8,)
|
|
@@ -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]
|
|
@@ -6,7 +6,7 @@ from functools import partial, cache
|
|
|
6
6
|
from itertools import zip_longest
|
|
7
7
|
|
|
8
8
|
import torch
|
|
9
|
-
from torch import nn, Tensor
|
|
9
|
+
from torch import nn, Tensor, arange, cat
|
|
10
10
|
from torch.nn import Module, ModuleList
|
|
11
11
|
import torch.nn.functional as F
|
|
12
12
|
import torch.distributed as dist
|
|
@@ -14,6 +14,7 @@ from vector_quantize_pytorch.vector_quantize_pytorch import VectorQuantize
|
|
|
14
14
|
|
|
15
15
|
from einops import rearrange, repeat, reduce, pack, unpack
|
|
16
16
|
|
|
17
|
+
import einx
|
|
17
18
|
from einx import get_at
|
|
18
19
|
|
|
19
20
|
# helper functions
|
|
@@ -36,6 +37,47 @@ def unique(arr):
|
|
|
36
37
|
def round_up_multiple(num, mult):
|
|
37
38
|
return ceil(num / mult) * mult
|
|
38
39
|
|
|
40
|
+
# tensor helpers
|
|
41
|
+
|
|
42
|
+
def pad_at_dim(
|
|
43
|
+
t,
|
|
44
|
+
pad: tuple[int, int],
|
|
45
|
+
dim = -1,
|
|
46
|
+
value = 0.
|
|
47
|
+
):
|
|
48
|
+
if pad == (0, 0):
|
|
49
|
+
return t
|
|
50
|
+
|
|
51
|
+
dims_from_right = (- dim - 1) if dim < 0 else (t.ndim - dim - 1)
|
|
52
|
+
zeros = ((0, 0) * dims_from_right)
|
|
53
|
+
return F.pad(t, (*zeros, *pad), value = value)
|
|
54
|
+
|
|
55
|
+
def pack_one(t, pattern):
|
|
56
|
+
packed, packed_shape = pack([t], pattern)
|
|
57
|
+
|
|
58
|
+
def inverse(out, inv_pattern = None):
|
|
59
|
+
inv_pattern = default(inv_pattern, pattern)
|
|
60
|
+
return first(unpack(out, packed_shape, inv_pattern))
|
|
61
|
+
|
|
62
|
+
return packed, inverse
|
|
63
|
+
|
|
64
|
+
def batch_select(t, indices, pattern = None):
|
|
65
|
+
|
|
66
|
+
if exists(pattern):
|
|
67
|
+
indices = rearrange(indices, '... k -> (...) k')
|
|
68
|
+
t, inv_pack = pack_one(t, pattern)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
batch_indices = arange(t.shape[0], device = t.device)
|
|
72
|
+
batch_indices = rearrange(batch_indices, 'b -> b 1')
|
|
73
|
+
|
|
74
|
+
out = t[batch_indices, indices]
|
|
75
|
+
|
|
76
|
+
if exists(pattern):
|
|
77
|
+
out = inv_pack(out)
|
|
78
|
+
|
|
79
|
+
return out
|
|
80
|
+
|
|
39
81
|
# distributed helpers
|
|
40
82
|
|
|
41
83
|
def is_distributed():
|
|
@@ -128,6 +170,8 @@ class ResidualVQ(Module):
|
|
|
128
170
|
accept_image_fmap = False,
|
|
129
171
|
implicit_neural_codebook = False, # QINCo from https://arxiv.org/abs/2401.14732
|
|
130
172
|
mlp_kwargs: dict = dict(),
|
|
173
|
+
beam_size = None,
|
|
174
|
+
eval_beam_size = None,
|
|
131
175
|
**vq_kwargs
|
|
132
176
|
):
|
|
133
177
|
super().__init__()
|
|
@@ -183,6 +227,15 @@ class ResidualVQ(Module):
|
|
|
183
227
|
self.quantize_dropout_cutoff_index = quantize_dropout_cutoff_index
|
|
184
228
|
self.quantize_dropout_multiple_of = quantize_dropout_multiple_of # encodec paper proposes structured dropout, believe this was set to 4
|
|
185
229
|
|
|
230
|
+
# determine whether is using ema update
|
|
231
|
+
|
|
232
|
+
self.vq_is_ema_updating = first(self.layers).ema_update
|
|
233
|
+
|
|
234
|
+
# beam size
|
|
235
|
+
|
|
236
|
+
self.beam_size = default(beam_size, eval_beam_size)
|
|
237
|
+
self.eval_beam_size = eval_beam_size
|
|
238
|
+
|
|
186
239
|
# setting up the MLPs for implicit neural codebooks
|
|
187
240
|
|
|
188
241
|
self.mlps = None
|
|
@@ -295,19 +348,29 @@ class ResidualVQ(Module):
|
|
|
295
348
|
return_all_codes = False,
|
|
296
349
|
sample_codebook_temp = None,
|
|
297
350
|
freeze_codebook = False,
|
|
351
|
+
beam_size = None,
|
|
298
352
|
rand_quantize_dropout_fixed_seed = None
|
|
299
353
|
):
|
|
300
|
-
|
|
354
|
+
|
|
355
|
+
# variables
|
|
356
|
+
|
|
357
|
+
input_shape, num_quant, quant_dropout_multiple_of, return_loss, device = x.shape, self.num_quantizers, self.quantize_dropout_multiple_of, exists(indices), x.device
|
|
358
|
+
|
|
359
|
+
beam_size = default(beam_size, self.beam_size if self.training else self.eval_beam_size)
|
|
360
|
+
|
|
361
|
+
is_beam_search = exists(beam_size) and beam_size > 1
|
|
362
|
+
|
|
363
|
+
# projecting in
|
|
301
364
|
|
|
302
365
|
x = self.project_in(x)
|
|
303
366
|
|
|
304
367
|
assert not (self.accept_image_fmap and exists(indices))
|
|
305
368
|
|
|
306
|
-
quantized_out =
|
|
369
|
+
quantized_out = torch.zeros_like(x)
|
|
307
370
|
residual = x
|
|
308
371
|
|
|
309
|
-
all_losses =
|
|
310
|
-
all_indices = []
|
|
372
|
+
all_losses = torch.empty((0,), device = device, dtype = torch.float32)
|
|
373
|
+
all_indices = torch.empty((*input_shape[:-1], 0), device = device, dtype = torch.long)
|
|
311
374
|
|
|
312
375
|
if isinstance(indices, list):
|
|
313
376
|
indices = torch.stack(indices)
|
|
@@ -319,7 +382,6 @@ class ResidualVQ(Module):
|
|
|
319
382
|
should_quantize_dropout = self.training and self.quantize_dropout and not return_loss
|
|
320
383
|
|
|
321
384
|
# sample a layer index at which to dropout further residual quantization
|
|
322
|
-
# also prepare null indices and loss
|
|
323
385
|
|
|
324
386
|
if should_quantize_dropout:
|
|
325
387
|
|
|
@@ -335,9 +397,24 @@ class ResidualVQ(Module):
|
|
|
335
397
|
if quant_dropout_multiple_of != 1:
|
|
336
398
|
rand_quantize_dropout_index = round_up_multiple(rand_quantize_dropout_index + 1, quant_dropout_multiple_of) - 1
|
|
337
399
|
|
|
338
|
-
|
|
339
|
-
|
|
340
|
-
|
|
400
|
+
# save all inputs across layers, for use during expiration at end under shared codebook setting, or ema update during beam search
|
|
401
|
+
|
|
402
|
+
all_residuals = torch.empty((*input_shape[:-1], 0, input_shape[-1]), dtype = residual.dtype, device = device)
|
|
403
|
+
|
|
404
|
+
# maybe prepare beam search
|
|
405
|
+
|
|
406
|
+
if is_beam_search:
|
|
407
|
+
prec_dims = x.shape[:-1]
|
|
408
|
+
|
|
409
|
+
search_scores = torch.zeros((*prec_dims, 1), device = device, dtype = x.dtype)
|
|
410
|
+
|
|
411
|
+
residual = rearrange(residual, '... d -> ... 1 d')
|
|
412
|
+
quantized_out = rearrange(quantized_out, '... d -> ... 1 d')
|
|
413
|
+
|
|
414
|
+
all_residuals = rearrange(all_residuals, '... l d -> ... 1 l d')
|
|
415
|
+
all_indices = rearrange(all_indices, '... l -> ... 1 l')
|
|
416
|
+
|
|
417
|
+
all_losses = all_losses.reshape(*input_shape[:-1], 1, 0)
|
|
341
418
|
|
|
342
419
|
# setup the mlps for implicit neural codebook
|
|
343
420
|
|
|
@@ -346,17 +423,13 @@ class ResidualVQ(Module):
|
|
|
346
423
|
if self.implicit_neural_codebook:
|
|
347
424
|
maybe_code_transforms = (None, *self.mlps)
|
|
348
425
|
|
|
349
|
-
# save all inputs across layers, for use during expiration at end under shared codebook setting
|
|
350
|
-
|
|
351
|
-
all_residuals = []
|
|
352
|
-
|
|
353
426
|
# go through the layers
|
|
354
427
|
|
|
355
428
|
for quantizer_index, (vq, maybe_mlp) in enumerate(zip(self.layers, maybe_code_transforms)):
|
|
356
429
|
|
|
357
430
|
if should_quantize_dropout and quantizer_index > rand_quantize_dropout_index:
|
|
358
|
-
all_indices
|
|
359
|
-
all_losses
|
|
431
|
+
all_indices = pad_at_dim(all_indices, (0, 1), value = -1, dim = -1)
|
|
432
|
+
all_losses = pad_at_dim(all_losses, (0, 1), value = 0, dim = -1)
|
|
360
433
|
continue
|
|
361
434
|
|
|
362
435
|
layer_indices = None
|
|
@@ -368,9 +441,9 @@ class ResidualVQ(Module):
|
|
|
368
441
|
if exists(maybe_mlp):
|
|
369
442
|
maybe_mlp = partial(maybe_mlp, condition = quantized_out)
|
|
370
443
|
|
|
371
|
-
# save for expiration
|
|
444
|
+
# save the residual input for maybe expiration as well as ema update after beam search
|
|
372
445
|
|
|
373
|
-
all_residuals
|
|
446
|
+
all_residuals = cat((all_residuals, rearrange(residual, '... d -> ... 1 d')), dim = -2)
|
|
374
447
|
|
|
375
448
|
# vector quantize forward
|
|
376
449
|
|
|
@@ -380,29 +453,111 @@ class ResidualVQ(Module):
|
|
|
380
453
|
indices = layer_indices,
|
|
381
454
|
sample_codebook_temp = sample_codebook_temp,
|
|
382
455
|
freeze_codebook = freeze_codebook,
|
|
383
|
-
codebook_transform_fn = maybe_mlp
|
|
456
|
+
codebook_transform_fn = maybe_mlp,
|
|
457
|
+
topk = beam_size if is_beam_search else None
|
|
384
458
|
)
|
|
385
459
|
|
|
386
|
-
|
|
387
|
-
quantized_out = quantized_out + quantized
|
|
460
|
+
# cross entropy loss for some old paper
|
|
388
461
|
|
|
389
462
|
if return_loss:
|
|
390
|
-
ce_loss = rest
|
|
463
|
+
ce_loss = first(rest)
|
|
391
464
|
ce_losses.append(ce_loss)
|
|
392
465
|
continue
|
|
393
466
|
|
|
394
467
|
embed_indices, loss = rest
|
|
395
468
|
|
|
396
|
-
|
|
397
|
-
|
|
469
|
+
# handle expanding first residual if doing beam search
|
|
470
|
+
|
|
471
|
+
if is_beam_search:
|
|
472
|
+
|
|
473
|
+
search_scores = einx.add('... j, ... j k -> ... (j k)', search_scores, -loss)
|
|
474
|
+
|
|
475
|
+
residual = rearrange(residual, '... j d -> ... j 1 d')
|
|
476
|
+
quantized_out = rearrange(quantized_out, '... j d -> ... j 1 d')
|
|
477
|
+
|
|
478
|
+
all_residuals = repeat(all_residuals, '... j l d -> ... (j k) l d', k = beam_size)
|
|
479
|
+
|
|
480
|
+
# core residual vq logic
|
|
481
|
+
|
|
482
|
+
residual = residual - quantized.detach()
|
|
483
|
+
quantized_out = quantized_out + quantized
|
|
484
|
+
|
|
485
|
+
# handle sort and topk beams
|
|
486
|
+
|
|
487
|
+
if is_beam_search:
|
|
488
|
+
residual = rearrange(residual, '... j k d -> ... (j k) d')
|
|
489
|
+
quantized_out = rearrange(quantized_out, '... j k d -> ... (j k) d')
|
|
490
|
+
|
|
491
|
+
# broadcat the indices
|
|
492
|
+
|
|
493
|
+
all_indices = repeat(all_indices, '... j l -> ... j k l', k = embed_indices.shape[-1])
|
|
494
|
+
embed_indices = rearrange(embed_indices, '... j k -> ... j k 1')
|
|
495
|
+
|
|
496
|
+
all_indices = cat((all_indices, embed_indices), dim = -1)
|
|
497
|
+
all_indices = rearrange(all_indices, '... j k l -> ... (j k) l')
|
|
498
|
+
|
|
499
|
+
# broadcat the losses
|
|
500
|
+
|
|
501
|
+
all_losses = repeat(all_losses, '... j l -> ... j k l', k = loss.shape[-1])
|
|
502
|
+
loss = rearrange(loss, '... -> ... 1')
|
|
503
|
+
|
|
504
|
+
all_losses = cat((all_losses, loss), dim = -1)
|
|
505
|
+
all_losses = rearrange(all_losses, '... j k l -> ... (j k) l')
|
|
506
|
+
|
|
507
|
+
# handle sort and selection of highest beam size
|
|
508
|
+
|
|
509
|
+
if search_scores.shape[-1] > beam_size:
|
|
510
|
+
search_scores, select_indices = search_scores.topk(beam_size, dim = -1)
|
|
511
|
+
|
|
512
|
+
residual = batch_select(residual, select_indices, '* k d')
|
|
513
|
+
quantized_out = batch_select(quantized_out, select_indices, '* k d')
|
|
514
|
+
|
|
515
|
+
all_indices = batch_select(all_indices, select_indices, '* k l')
|
|
516
|
+
all_losses = batch_select(all_losses, select_indices, '* k l')
|
|
517
|
+
|
|
518
|
+
all_residuals = batch_select(all_residuals, select_indices, '* k l d')
|
|
519
|
+
else:
|
|
520
|
+
# aggregate indices and losses
|
|
521
|
+
|
|
522
|
+
all_indices = cat((all_indices, rearrange(embed_indices, '... -> ... 1')), dim = -1)
|
|
523
|
+
|
|
524
|
+
all_losses = cat((all_losses, rearrange(loss, '... -> ... 1')), dim = -1)
|
|
525
|
+
|
|
526
|
+
# handle beam search
|
|
527
|
+
|
|
528
|
+
if is_beam_search:
|
|
529
|
+
top_index = search_scores.argmax(dim = -1, keepdim = True)
|
|
530
|
+
|
|
531
|
+
quantized_out = batch_select(quantized_out, top_index, '* k d')
|
|
532
|
+
all_indices = batch_select(all_indices, top_index, '* k l')
|
|
533
|
+
all_losses = batch_select(all_losses, top_index, '* k l')
|
|
534
|
+
all_residuals = batch_select(all_residuals, top_index, '* k l d')
|
|
535
|
+
|
|
536
|
+
quantized_out, all_indices, all_losses, all_residuals = [t[..., 0, :] for t in (quantized_out, all_indices, all_losses, all_residuals)]
|
|
537
|
+
|
|
538
|
+
# handle commit loss, which should be the average
|
|
539
|
+
|
|
540
|
+
if exists(mask):
|
|
541
|
+
all_losses = einx.where('..., ... l,', mask, all_losses, 0.)
|
|
542
|
+
all_losses = reduce(all_losses, '... l -> l', 'sum') / mask.sum(dim = -1).clamp_min(1e-4)
|
|
543
|
+
else:
|
|
544
|
+
all_losses = reduce(all_losses, '... l -> l', 'mean')
|
|
545
|
+
|
|
546
|
+
# handle updating ema
|
|
547
|
+
|
|
548
|
+
if self.vq_is_ema_updating:
|
|
549
|
+
for vq, layer_input, indices in zip(self.layers, all_residuals.unbind(dim = -2), all_indices.unbind(dim = -1)): # in the case of quantize dropout, zip will terminate with the shorter sequence, which should be all_residuals
|
|
550
|
+
vq.update_ema_indices(layer_input, indices, mask = mask)
|
|
398
551
|
|
|
399
552
|
# if shared codebook, update ema only at end
|
|
400
553
|
|
|
401
|
-
if self.training and self.shared_codebook:
|
|
554
|
+
if self.training and self.shared_codebook and not is_beam_search:
|
|
402
555
|
shared_layer = first(self.layers)
|
|
403
556
|
shared_layer._codebook.update_ema()
|
|
404
557
|
shared_layer.update_in_place_optimizer()
|
|
405
|
-
|
|
558
|
+
|
|
559
|
+
all_codes_for_expire = rearrange(all_residuals, '... n l d -> ... (n l) d')
|
|
560
|
+
shared_layer.expire_codes_(all_codes_for_expire)
|
|
406
561
|
|
|
407
562
|
# project out, if needed
|
|
408
563
|
|
|
@@ -415,8 +570,6 @@ class ResidualVQ(Module):
|
|
|
415
570
|
|
|
416
571
|
# stack all losses and indices
|
|
417
572
|
|
|
418
|
-
all_losses, all_indices = map(partial(torch.stack, dim = -1), (all_losses, all_indices))
|
|
419
|
-
|
|
420
573
|
ret = (quantized_out, all_indices, all_losses)
|
|
421
574
|
|
|
422
575
|
if return_all_codes:
|