vector-quantize-pytorch 1.31.0__tar.gz → 1.31.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.
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/PKG-INFO +3 -2
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/pyproject.toml +3 -2
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/__init__.py +1 -0
- vector_quantize_pytorch-1.31.2/vector_quantize_pytorch/evo_lookup_free_quantization.py +289 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +57 -30
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/.gitignore +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/LICENSE +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/README.md +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/finite_scalar_perturbation.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/hierarchical_vq.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/residual_vq.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/sim_vq.py +0 -0
- {vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
2
|
Name: vector-quantize-pytorch
|
|
3
|
-
Version: 1.31.
|
|
3
|
+
Version: 1.31.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
|
|
@@ -36,6 +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-einops-utils>=0.1.12
|
|
39
40
|
Requires-Dist: torch>=2.4
|
|
40
41
|
Provides-Extra: examples
|
|
41
42
|
Requires-Dist: torchvision; extra == 'examples'
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "vector-quantize-pytorch"
|
|
3
|
-
version = "1.31.
|
|
3
|
+
version = "1.31.2"
|
|
4
4
|
description = "Vector Quantization - Pytorch"
|
|
5
5
|
authors = [
|
|
6
6
|
{ name = "Phil Wang", email = "lucidrains@gmail.com" }
|
|
@@ -23,9 +23,10 @@ classifiers=[
|
|
|
23
23
|
]
|
|
24
24
|
|
|
25
25
|
dependencies = [
|
|
26
|
-
"torch>=2.4",
|
|
27
26
|
"einops>=0.8.0",
|
|
28
27
|
"einx>=0.3.0",
|
|
28
|
+
"torch>=2.4",
|
|
29
|
+
"torch-einops-utils>=0.1.12",
|
|
29
30
|
]
|
|
30
31
|
|
|
31
32
|
[project.urls]
|
|
@@ -4,6 +4,7 @@ from vector_quantize_pytorch.random_projection_quantizer import RandomProjection
|
|
|
4
4
|
from vector_quantize_pytorch.finite_scalar_quantization import FSQ
|
|
5
5
|
from vector_quantize_pytorch.finite_scalar_perturbation import FSP
|
|
6
6
|
from vector_quantize_pytorch.lookup_free_quantization import LFQ
|
|
7
|
+
from vector_quantize_pytorch.evo_lookup_free_quantization import EvoLFQ
|
|
7
8
|
from vector_quantize_pytorch.residual_lfq import ResidualLFQ, GroupedResidualLFQ
|
|
8
9
|
from vector_quantize_pytorch.residual_fsq import ResidualFSQ, GroupedResidualFSQ
|
|
9
10
|
from vector_quantize_pytorch.latent_quantization import LatentQuantize
|
|
@@ -0,0 +1,289 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
from collections import namedtuple
|
|
3
|
+
|
|
4
|
+
import torch
|
|
5
|
+
from torch import nn, Tensor, is_tensor
|
|
6
|
+
from torch.nn import Module
|
|
7
|
+
import torch.nn.functional as F
|
|
8
|
+
|
|
9
|
+
from einops import rearrange, repeat, reduce
|
|
10
|
+
from torch_einops_utils import pack_with_inverse, temp_eval
|
|
11
|
+
|
|
12
|
+
from vector_quantize_pytorch.lookup_free_quantization import LFQ
|
|
13
|
+
|
|
14
|
+
# constants
|
|
15
|
+
|
|
16
|
+
Return = namedtuple('Return', ['reconstructed', 'indices', 'entropy_aux_loss'])
|
|
17
|
+
Result = namedtuple('Result', ['pop_bits', 'best_gene', 'best_fitness', 'best_decoded'])
|
|
18
|
+
|
|
19
|
+
# helper functions
|
|
20
|
+
|
|
21
|
+
def exists(val):
|
|
22
|
+
return val is not None
|
|
23
|
+
|
|
24
|
+
def default(val, d):
|
|
25
|
+
return val if exists(val) else d
|
|
26
|
+
|
|
27
|
+
# main class
|
|
28
|
+
|
|
29
|
+
class EvoLFQ(Module):
|
|
30
|
+
def __init__(
|
|
31
|
+
self,
|
|
32
|
+
encoder: Module,
|
|
33
|
+
decoder: Module,
|
|
34
|
+
lfq: LFQ | None = None,
|
|
35
|
+
*,
|
|
36
|
+
dim: int | None = None,
|
|
37
|
+
codebook_size: int | None = None,
|
|
38
|
+
num_codebooks: int = 1,
|
|
39
|
+
pop_size: int = 64,
|
|
40
|
+
mutation_rate: float = 0.02,
|
|
41
|
+
tournament_size: int = 2,
|
|
42
|
+
elitism_count: int = 1,
|
|
43
|
+
generations: int = 50,
|
|
44
|
+
**lfq_kwargs
|
|
45
|
+
):
|
|
46
|
+
super().__init__()
|
|
47
|
+
assert pop_size > elitism_count, 'pop_size must be greater than elitism_count'
|
|
48
|
+
|
|
49
|
+
self.encoder = encoder
|
|
50
|
+
self.decoder = decoder
|
|
51
|
+
|
|
52
|
+
if not exists(lfq):
|
|
53
|
+
assert exists(dim) or exists(codebook_size), 'either lfq instance or dim / codebook_size must be supplied to EvoLFQ'
|
|
54
|
+
lfq = LFQ(
|
|
55
|
+
dim = dim,
|
|
56
|
+
codebook_size = codebook_size,
|
|
57
|
+
num_codebooks = num_codebooks,
|
|
58
|
+
**lfq_kwargs
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
self.lfq = lfq
|
|
62
|
+
self.pop_size = pop_size
|
|
63
|
+
self.mutation_rate = mutation_rate
|
|
64
|
+
self.tournament_size = tournament_size
|
|
65
|
+
self.elitism_count = elitism_count
|
|
66
|
+
self.generations = generations
|
|
67
|
+
|
|
68
|
+
self.register_buffer('zero', torch.tensor(0.), persistent = False)
|
|
69
|
+
|
|
70
|
+
@property
|
|
71
|
+
def device(self):
|
|
72
|
+
return self.zero.device
|
|
73
|
+
|
|
74
|
+
def forward(self, x, **kwargs):
|
|
75
|
+
latents = self.encoder(x)
|
|
76
|
+
is_2d = latents.ndim == 2
|
|
77
|
+
|
|
78
|
+
if is_2d:
|
|
79
|
+
latents = rearrange(latents, 'b d -> b 1 d')
|
|
80
|
+
|
|
81
|
+
quantized, indices, aux_loss = self.lfq(latents, **kwargs)
|
|
82
|
+
|
|
83
|
+
if is_2d:
|
|
84
|
+
quantized = rearrange(quantized, 'b 1 d -> b d')
|
|
85
|
+
indices = rearrange(indices, 'b 1 ... -> b ...')
|
|
86
|
+
|
|
87
|
+
reconstructed = self.decoder(quantized)
|
|
88
|
+
return Return(reconstructed, indices, aux_loss)
|
|
89
|
+
|
|
90
|
+
@torch.no_grad()
|
|
91
|
+
def encode(self, x, return_signs = False):
|
|
92
|
+
with temp_eval(self):
|
|
93
|
+
latents = self.encoder(x)
|
|
94
|
+
is_2d = latents.ndim == 2
|
|
95
|
+
|
|
96
|
+
if is_2d:
|
|
97
|
+
latents = rearrange(latents, 'b d -> b 1 d')
|
|
98
|
+
|
|
99
|
+
quantized, indices, _ = self.lfq(latents)
|
|
100
|
+
|
|
101
|
+
if is_2d:
|
|
102
|
+
quantized = rearrange(quantized, 'b 1 d -> b d')
|
|
103
|
+
|
|
104
|
+
if return_signs:
|
|
105
|
+
return torch.where(quantized > 0, 1.0, -1.0)
|
|
106
|
+
|
|
107
|
+
return (quantized > 0).float()
|
|
108
|
+
|
|
109
|
+
@torch.no_grad()
|
|
110
|
+
def decode_bits(self, bits):
|
|
111
|
+
"""
|
|
112
|
+
Converts binary bits (0/1 or -1/+1) to codebook indices and uses lfq.indices_to_codes
|
|
113
|
+
to decode back to data space.
|
|
114
|
+
"""
|
|
115
|
+
with temp_eval(self):
|
|
116
|
+
bits = (bits > 0).int()
|
|
117
|
+
bits, inverse = pack_with_inverse(bits, '* d')
|
|
118
|
+
|
|
119
|
+
codebook_dim = self.lfq.codebook_dim
|
|
120
|
+
num_codebooks = self.lfq.num_codebooks
|
|
121
|
+
|
|
122
|
+
if num_codebooks > 1 and bits.shape[-1] == codebook_dim * num_codebooks:
|
|
123
|
+
bits = rearrange(bits, 'b (c d) -> b c d', c = num_codebooks, d = codebook_dim)
|
|
124
|
+
|
|
125
|
+
if bits.ndim == 2:
|
|
126
|
+
bits = rearrange(bits, 'b d -> b 1 d')
|
|
127
|
+
|
|
128
|
+
mask = 2 ** torch.arange(codebook_dim - 1, -1, -1, device = self.device)
|
|
129
|
+
indices = reduce(bits * mask, '... d -> ...', 'sum')
|
|
130
|
+
|
|
131
|
+
if not self.lfq.keep_num_codebooks_dim and indices.ndim >= 2 and indices.shape[-1] == 1:
|
|
132
|
+
indices = rearrange(indices, '... 1 -> ...')
|
|
133
|
+
|
|
134
|
+
codes = self.lfq.indices_to_codes(indices)
|
|
135
|
+
|
|
136
|
+
if codes.ndim == 3 and codes.shape[1] == 1:
|
|
137
|
+
codes = rearrange(codes, 'b 1 d -> b d')
|
|
138
|
+
|
|
139
|
+
codes = inverse(codes, '* d')
|
|
140
|
+
return self.decoder(codes)
|
|
141
|
+
|
|
142
|
+
# genetic algorithm helpers
|
|
143
|
+
|
|
144
|
+
def init_random_population(self, pop_size = None, shape = None, device = None, is_sign = False):
|
|
145
|
+
pop_size = default(pop_size, self.pop_size)
|
|
146
|
+
device = default(device, self.device)
|
|
147
|
+
bits = torch.randint(0, 2, (pop_size, *shape), device = device).float()
|
|
148
|
+
|
|
149
|
+
if is_sign:
|
|
150
|
+
return torch.where(bits > 0, 1.0, -1.0)
|
|
151
|
+
|
|
152
|
+
return bits
|
|
153
|
+
|
|
154
|
+
def init_population_from_data(self, x, pop_size = None, mutation_rate = None, is_sign = False):
|
|
155
|
+
pop_size = default(pop_size, self.pop_size)
|
|
156
|
+
mutation_rate = default(mutation_rate, self.mutation_rate)
|
|
157
|
+
bits = self.encode(x, return_signs = is_sign)
|
|
158
|
+
|
|
159
|
+
num_samples, *_ = bits.shape
|
|
160
|
+
|
|
161
|
+
if num_samples < pop_size:
|
|
162
|
+
repeats = (pop_size + num_samples - 1) // num_samples
|
|
163
|
+
bits = repeat(bits, 'b ... -> (r b) ...', r = repeats)[:pop_size]
|
|
164
|
+
else:
|
|
165
|
+
bits = bits[:pop_size]
|
|
166
|
+
|
|
167
|
+
if mutation_rate > 0:
|
|
168
|
+
bits = self.mutate(bits, mutation_rate = mutation_rate, is_sign = is_sign)
|
|
169
|
+
|
|
170
|
+
return bits
|
|
171
|
+
|
|
172
|
+
def uniform_crossover(self, parent1, parent2):
|
|
173
|
+
crossover_mask = torch.rand_like(parent1.float()) < 0.5
|
|
174
|
+
return torch.where(crossover_mask, parent1, parent2)
|
|
175
|
+
|
|
176
|
+
def mutate(self, population, mutation_rate = None, is_sign = False):
|
|
177
|
+
mutation_rate = default(mutation_rate, self.mutation_rate)
|
|
178
|
+
flip_mask = torch.rand_like(population.float()) < mutation_rate
|
|
179
|
+
|
|
180
|
+
if is_sign:
|
|
181
|
+
return torch.where(flip_mask, -population, population)
|
|
182
|
+
|
|
183
|
+
return torch.where(flip_mask, 1.0 - population, population)
|
|
184
|
+
|
|
185
|
+
def tournament_selection(self, population, fitnesses, tournament_size = None):
|
|
186
|
+
tournament_size = default(tournament_size, self.tournament_size)
|
|
187
|
+
pop_size, *_ = population.shape
|
|
188
|
+
|
|
189
|
+
contenders = torch.randint(0, pop_size, (pop_size, tournament_size), device = self.device)
|
|
190
|
+
winner_indices = contenders.gather(1, fitnesses[contenders].argmax(dim = -1, keepdim = True))
|
|
191
|
+
return population[rearrange(winner_indices, 'b 1 -> b')]
|
|
192
|
+
|
|
193
|
+
@torch.no_grad()
|
|
194
|
+
def step(
|
|
195
|
+
self,
|
|
196
|
+
pop_bits,
|
|
197
|
+
fitness_fn,
|
|
198
|
+
tournament_size = None,
|
|
199
|
+
mutation_rate = None,
|
|
200
|
+
elitism_count = None,
|
|
201
|
+
is_sign = False,
|
|
202
|
+
batch_size = None,
|
|
203
|
+
**kwargs
|
|
204
|
+
):
|
|
205
|
+
tournament_size = default(tournament_size, self.tournament_size)
|
|
206
|
+
mutation_rate = default(mutation_rate, self.mutation_rate)
|
|
207
|
+
elitism_count = default(elitism_count, self.elitism_count)
|
|
208
|
+
|
|
209
|
+
pop_size, *_ = pop_bits.shape
|
|
210
|
+
num_offspring = pop_size - elitism_count
|
|
211
|
+
|
|
212
|
+
# evaluate fitness
|
|
213
|
+
|
|
214
|
+
if exists(batch_size):
|
|
215
|
+
decoded_chunks = [self.decode_bits(chunk) for chunk in pop_bits.split(batch_size)]
|
|
216
|
+
decoded = torch.cat(decoded_chunks, dim = 0)
|
|
217
|
+
else:
|
|
218
|
+
decoded = self.decode_bits(pop_bits)
|
|
219
|
+
|
|
220
|
+
fitnesses = fitness_fn(decoded, pop_bits)
|
|
221
|
+
if not is_tensor(fitnesses):
|
|
222
|
+
fitnesses = torch.tensor(fitnesses, device = self.device, dtype = torch.float32)
|
|
223
|
+
|
|
224
|
+
# sort by fitness descending and preserve elites
|
|
225
|
+
|
|
226
|
+
sorted_indices = torch.argsort(fitnesses, descending = True)
|
|
227
|
+
elites = pop_bits[sorted_indices[:elitism_count]].clone()
|
|
228
|
+
|
|
229
|
+
# tournament selection: 2 parent tournaments for each of the num_offspring children needed
|
|
230
|
+
|
|
231
|
+
contenders = torch.randint(0, pop_size, (num_offspring, 2, tournament_size), device = self.device)
|
|
232
|
+
winner_rel_idx = fitnesses[contenders].argmax(dim = -1, keepdim = True)
|
|
233
|
+
parent_indices = rearrange(contenders.gather(-1, winner_rel_idx), 'n p 1 -> n p')
|
|
234
|
+
|
|
235
|
+
# crossover: 1 child produced per 2 parents
|
|
236
|
+
|
|
237
|
+
p1, p2 = rearrange(pop_bits[parent_indices], 'n p ... -> p n ...')
|
|
238
|
+
offspring = self.uniform_crossover(p1, p2)
|
|
239
|
+
|
|
240
|
+
# mutation all at once
|
|
241
|
+
|
|
242
|
+
offspring = self.mutate(offspring, mutation_rate = mutation_rate, is_sign = is_sign)
|
|
243
|
+
|
|
244
|
+
next_pop = torch.cat([elites, offspring], dim = 0)
|
|
245
|
+
return next_pop, fitnesses
|
|
246
|
+
|
|
247
|
+
@torch.no_grad()
|
|
248
|
+
def evolve(
|
|
249
|
+
self,
|
|
250
|
+
fitness_fn,
|
|
251
|
+
pop_bits = None,
|
|
252
|
+
pop_size = None,
|
|
253
|
+
shape = None,
|
|
254
|
+
generations = None,
|
|
255
|
+
is_sign = False,
|
|
256
|
+
return_best_decoded = False,
|
|
257
|
+
**step_kwargs
|
|
258
|
+
):
|
|
259
|
+
pop_size = default(pop_size, self.pop_size)
|
|
260
|
+
generations = default(generations, self.generations)
|
|
261
|
+
|
|
262
|
+
if not exists(pop_bits):
|
|
263
|
+
assert exists(shape), 'shape must be provided if pop_bits is not supplied to evolve()'
|
|
264
|
+
pop_bits = self.init_random_population(pop_size, shape, is_sign = is_sign)
|
|
265
|
+
|
|
266
|
+
best_fitness = float('-inf')
|
|
267
|
+
best_gene = None
|
|
268
|
+
|
|
269
|
+
for _ in range(generations):
|
|
270
|
+
pop_bits, fitnesses = self.step(
|
|
271
|
+
pop_bits,
|
|
272
|
+
fitness_fn,
|
|
273
|
+
is_sign = is_sign,
|
|
274
|
+
**step_kwargs
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
max_fit_idx = torch.argmax(fitnesses)
|
|
278
|
+
max_fit = fitnesses[max_fit_idx].item()
|
|
279
|
+
|
|
280
|
+
if max_fit > best_fitness:
|
|
281
|
+
best_fitness = max_fit
|
|
282
|
+
best_gene = pop_bits[max_fit_idx].clone()
|
|
283
|
+
|
|
284
|
+
best_decoded = None
|
|
285
|
+
if return_best_decoded:
|
|
286
|
+
best_decoded = self.decode_bits(rearrange(best_gene, '... -> 1 ...'))
|
|
287
|
+
best_decoded = rearrange(best_decoded, '1 ... -> ...')
|
|
288
|
+
|
|
289
|
+
yield Result(pop_bits, best_gene, best_fitness, best_decoded)
|
|
@@ -654,14 +654,15 @@ class Codebook(Module):
|
|
|
654
654
|
|
|
655
655
|
if needs_codebook_dim:
|
|
656
656
|
x = rearrange(x, '... -> 1 ...')
|
|
657
|
+
embed_ind = rearrange(embed_ind, '... -> 1 ...')
|
|
657
658
|
|
|
658
659
|
dtype = x.dtype
|
|
659
|
-
flatten,
|
|
660
|
+
flatten, _ = pack_one(x, 'h * d')
|
|
660
661
|
|
|
661
662
|
if exists(mask):
|
|
662
663
|
mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
|
|
663
664
|
|
|
664
|
-
embed_ind, _ =
|
|
665
|
+
embed_ind, _ = pack_one(embed_ind, 'h *')
|
|
665
666
|
embed_ind = embed_ind.masked_fill(embed_ind == -1, 0)
|
|
666
667
|
embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
|
|
667
668
|
|
|
@@ -697,6 +698,8 @@ class Codebook(Module):
|
|
|
697
698
|
dtype = x.dtype
|
|
698
699
|
flatten, unpack_one = pack_one(x, 'h * d')
|
|
699
700
|
|
|
701
|
+
is_topk = exists(topk)
|
|
702
|
+
|
|
700
703
|
if exists(mask):
|
|
701
704
|
mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
|
|
702
705
|
|
|
@@ -746,19 +749,13 @@ class Codebook(Module):
|
|
|
746
749
|
|
|
747
750
|
embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, topk = topk, temperature = sample_codebook_temp, training = self.training)
|
|
748
751
|
|
|
749
|
-
if
|
|
750
|
-
embed_ind = unpack_one(embed_ind, 'h * k')
|
|
751
|
-
else:
|
|
752
|
-
embed_ind = unpack_one(embed_ind, 'h *')
|
|
752
|
+
embed_ind = unpack_one(embed_ind, 'h * k' if is_topk else 'h *')
|
|
753
753
|
|
|
754
754
|
if exists(codebook_transform_fn):
|
|
755
755
|
transformed_embed = unpack_one(transformed_embed, 'h * c d')
|
|
756
756
|
|
|
757
|
-
if self.training:
|
|
758
|
-
if
|
|
759
|
-
unpacked_onehot = unpack_one(embed_onehot, 'h * k c')
|
|
760
|
-
else:
|
|
761
|
-
unpacked_onehot = unpack_one(embed_onehot, 'h * c')
|
|
757
|
+
if self.training or is_topk:
|
|
758
|
+
unpacked_onehot = unpack_one(embed_onehot, 'h * k c' if is_topk else 'h * c')
|
|
762
759
|
|
|
763
760
|
if exists(codebook_transform_fn):
|
|
764
761
|
quantize = einsum('h b n ... c, h b n c d -> h b n ... d', unpacked_onehot, transformed_embed)
|
|
@@ -780,7 +777,7 @@ class Codebook(Module):
|
|
|
780
777
|
repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
|
|
781
778
|
quantize = repeated_embed.gather(-2, repeated_embed_ind)
|
|
782
779
|
|
|
783
|
-
if self.training and update_usage and not freeze_codebook and not
|
|
780
|
+
if self.training and update_usage and not freeze_codebook and not is_topk:
|
|
784
781
|
self.update_codebook(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update, ema_update = ema_update)
|
|
785
782
|
|
|
786
783
|
if needs_codebook_dim:
|
|
@@ -1107,6 +1104,18 @@ class VectorQuantize(Module):
|
|
|
1107
1104
|
):
|
|
1108
1105
|
orig_input, input_requires_grad = x, x.requires_grad
|
|
1109
1106
|
|
|
1107
|
+
is_topk = exists(topk)
|
|
1108
|
+
|
|
1109
|
+
if is_topk:
|
|
1110
|
+
topk = min(int(topk), self.codebook_size)
|
|
1111
|
+
|
|
1112
|
+
assert (
|
|
1113
|
+
self.heads == 1
|
|
1114
|
+
and topk > 0
|
|
1115
|
+
and not (self.commitment_use_cross_entropy_loss and self.training)
|
|
1116
|
+
and not (exists(self.in_place_codebook_optimizer) and self.training)
|
|
1117
|
+
), 'topk codes only supported for single-head codebooks without cross entropy commitment loss or in-place codebook optimizer'
|
|
1118
|
+
|
|
1110
1119
|
# freezing codebook
|
|
1111
1120
|
|
|
1112
1121
|
freeze_codebook = default(freeze_codebook, self.freeze_codebook)
|
|
@@ -1128,7 +1137,7 @@ class VectorQuantize(Module):
|
|
|
1128
1137
|
|
|
1129
1138
|
shape, dtype, device, heads, is_multiheaded, codebook_size, return_loss = x.shape, x.dtype, x.device, self.heads, self.heads > 1, self.codebook_size, exists(indices)
|
|
1130
1139
|
|
|
1131
|
-
need_transpose = not self.channel_last and not self.accept_image_fmap and not self.accept_3d_fmap
|
|
1140
|
+
need_transpose = not self.channel_last and not self.accept_image_fmap and not self.accept_3d_fmap and not only_one
|
|
1132
1141
|
should_inplace_optimize = exists(self.in_place_codebook_optimizer)
|
|
1133
1142
|
|
|
1134
1143
|
# rearrange inputs
|
|
@@ -1167,7 +1176,7 @@ class VectorQuantize(Module):
|
|
|
1167
1176
|
codebook_transform_fn = codebook_transform_fn,
|
|
1168
1177
|
ema_update_weight = ema_update_weight,
|
|
1169
1178
|
accum_ema_update = accum_ema_update,
|
|
1170
|
-
ema_update = ema_update and not
|
|
1179
|
+
ema_update = ema_update and not is_topk,
|
|
1171
1180
|
topk = topk
|
|
1172
1181
|
)
|
|
1173
1182
|
|
|
@@ -1209,17 +1218,14 @@ class VectorQuantize(Module):
|
|
|
1209
1218
|
codebook_forward_kwargs.update(update_usage = False)
|
|
1210
1219
|
quantize, embed_ind, distances = self._codebook(x, **codebook_forward_kwargs)
|
|
1211
1220
|
|
|
1221
|
+
if is_topk:
|
|
1222
|
+
x = repeat(x, '... d -> ... k d', k = topk)
|
|
1223
|
+
|
|
1212
1224
|
if self.training:
|
|
1213
|
-
# determine code to use for commitment loss
|
|
1214
1225
|
maybe_detach = torch.detach if not self.learnable_codebook or freeze_codebook else identity
|
|
1215
1226
|
|
|
1216
1227
|
commit_quantize = maybe_detach(quantize)
|
|
1217
1228
|
|
|
1218
|
-
# maybe expand input if returning topk codes
|
|
1219
|
-
|
|
1220
|
-
if exists(topk):
|
|
1221
|
-
x = repeat(x, '... d -> ... k d', k = topk)
|
|
1222
|
-
|
|
1223
1229
|
# spare rotation trick calculation if inputs do not need gradients
|
|
1224
1230
|
|
|
1225
1231
|
if input_requires_grad and self.route_gradients_to_input:
|
|
@@ -1279,7 +1285,19 @@ class VectorQuantize(Module):
|
|
|
1279
1285
|
|
|
1280
1286
|
# aggregate loss
|
|
1281
1287
|
|
|
1282
|
-
loss =
|
|
1288
|
+
loss = self.zero
|
|
1289
|
+
|
|
1290
|
+
if self.training:
|
|
1291
|
+
loss.requires_grad_()
|
|
1292
|
+
|
|
1293
|
+
if is_topk and not self.training:
|
|
1294
|
+
commit_loss = F.mse_loss(quantize, x, reduction = 'none')
|
|
1295
|
+
commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
|
|
1296
|
+
|
|
1297
|
+
if exists(mask):
|
|
1298
|
+
commit_loss = einx.where('..., ... k, -> ... k', mask, commit_loss, 0.)
|
|
1299
|
+
|
|
1300
|
+
loss = commit_loss * self.commitment_weight if self.has_commitment_loss else commit_loss
|
|
1283
1301
|
|
|
1284
1302
|
if self.training:
|
|
1285
1303
|
# calculate codebook diversity loss (negative of entropy) if needed
|
|
@@ -1304,11 +1322,8 @@ class VectorQuantize(Module):
|
|
|
1304
1322
|
|
|
1305
1323
|
commit_loss = calculate_ce_loss(embed_ind)
|
|
1306
1324
|
else:
|
|
1307
|
-
if
|
|
1308
|
-
|
|
1309
|
-
|
|
1310
|
-
repeated_input = repeat(orig_input, '... d -> ... k d', k = topk)
|
|
1311
|
-
commit_loss = F.mse_loss(commit_quantize, repeated_input, reduction = 'none')
|
|
1325
|
+
if is_topk:
|
|
1326
|
+
commit_loss = F.mse_loss(commit_quantize, x, reduction = 'none')
|
|
1312
1327
|
commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
|
|
1313
1328
|
|
|
1314
1329
|
if exists(mask):
|
|
@@ -1347,6 +1362,18 @@ class VectorQuantize(Module):
|
|
|
1347
1362
|
orthogonal_reg_loss = orthogonal_loss_fn(codebook)
|
|
1348
1363
|
loss = loss + orthogonal_reg_loss * self.orthogonal_reg_weight
|
|
1349
1364
|
|
|
1365
|
+
# reshape the topk loss to have the same shape as the embedding indices, if it is a per token loss
|
|
1366
|
+
|
|
1367
|
+
if is_topk and loss.ndim > 1:
|
|
1368
|
+
if only_one:
|
|
1369
|
+
loss = rearrange(loss, 'b 1 ... -> b ...')
|
|
1370
|
+
|
|
1371
|
+
if self.accept_image_fmap:
|
|
1372
|
+
loss = rearrange(loss, 'b (h w) ... -> b h w ...', h = height, w = width)
|
|
1373
|
+
|
|
1374
|
+
if self.accept_3d_fmap:
|
|
1375
|
+
loss = rearrange(loss, "b (d h w) ... -> b d h w ...", d = depth, h = height, w = width)
|
|
1376
|
+
|
|
1350
1377
|
# handle multi-headed quantized embeddings
|
|
1351
1378
|
|
|
1352
1379
|
if is_multiheaded:
|
|
@@ -1362,16 +1389,16 @@ class VectorQuantize(Module):
|
|
|
1362
1389
|
# rearrange quantized embeddings
|
|
1363
1390
|
|
|
1364
1391
|
if need_transpose:
|
|
1365
|
-
quantize = rearrange(quantize, 'b n d -> b d n')
|
|
1392
|
+
quantize = rearrange(quantize, 'b n ... d -> b d n ...')
|
|
1366
1393
|
|
|
1367
1394
|
if self.accept_image_fmap:
|
|
1368
|
-
quantize = rearrange(quantize, 'b (h w) c -> b c h w', h = height, w = width)
|
|
1395
|
+
quantize = rearrange(quantize, 'b (h w) ... c -> b c h w ...', h = height, w = width)
|
|
1369
1396
|
|
|
1370
1397
|
if self.accept_3d_fmap:
|
|
1371
|
-
quantize = rearrange(quantize, "b (d h w) c -> b c d h w", d=depth, h=height, w=width)
|
|
1398
|
+
quantize = rearrange(quantize, "b (d h w) ... c -> b c d h w ...", d=depth, h=height, w=width)
|
|
1372
1399
|
|
|
1373
1400
|
if only_one:
|
|
1374
|
-
quantize = rearrange(quantize, 'b 1 d -> b d')
|
|
1401
|
+
quantize = rearrange(quantize, 'b 1 ... d -> b ... d')
|
|
1375
1402
|
|
|
1376
1403
|
# if masking, only return quantized for where mask has True
|
|
1377
1404
|
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.31.0 → vector_quantize_pytorch-1.31.2}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|