vector-quantize-pytorch 1.26.0__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.26.0 → vector_quantize_pytorch-1.27.0}/PKG-INFO +1 -1
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/pyproject.toml +1 -1
- vector_quantize_pytorch-1.27.0/tests/test_beam.py +64 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_vq.py +180 -27
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/vector_quantize_pytorch.py +176 -278
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/build.yml +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/python-publish.yml +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/test.yml +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.gitignore +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/LICENSE +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/README.md +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_fsq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_lfq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/fsq.png +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/lfq.png +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/simvq.png +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/images/vq.png +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/ruff.toml +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/tests/test_latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/tests/test_lfq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/tests/test_readme.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/sim_vq.py +0 -0
- {vector_quantize_pytorch-1.26.0 → 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
|
|
@@ -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,)
|
|
@@ -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:
|
|
@@ -1,4 +1,5 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
|
+
from typing import Callable
|
|
2
3
|
|
|
3
4
|
from math import sqrt
|
|
4
5
|
from functools import partial, cache
|
|
@@ -6,7 +7,7 @@ from collections import namedtuple
|
|
|
6
7
|
|
|
7
8
|
import torch
|
|
8
9
|
from torch.nn import Module
|
|
9
|
-
from torch import nn, einsum, is_tensor, Tensor
|
|
10
|
+
from torch import nn, einsum, tensor, is_tensor, Tensor
|
|
10
11
|
import torch.nn.functional as F
|
|
11
12
|
import torch.distributed as distributed
|
|
12
13
|
from torch.optim import Optimizer
|
|
@@ -15,8 +16,6 @@ from torch.amp import autocast
|
|
|
15
16
|
import einx
|
|
16
17
|
from einops import rearrange, repeat, reduce, pack, unpack
|
|
17
18
|
|
|
18
|
-
from typing import Callable
|
|
19
|
-
|
|
20
19
|
def exists(val):
|
|
21
20
|
return val is not None
|
|
22
21
|
|
|
@@ -120,7 +119,8 @@ def gumbel_sample(
|
|
|
120
119
|
stochastic = False,
|
|
121
120
|
straight_through = False,
|
|
122
121
|
dim = -1,
|
|
123
|
-
training = True
|
|
122
|
+
training = True,
|
|
123
|
+
topk = None
|
|
124
124
|
):
|
|
125
125
|
dtype, size = logits.dtype, logits.shape[dim]
|
|
126
126
|
|
|
@@ -129,7 +129,11 @@ def gumbel_sample(
|
|
|
129
129
|
else:
|
|
130
130
|
sampling_logits = logits
|
|
131
131
|
|
|
132
|
-
|
|
132
|
+
if exists(topk):
|
|
133
|
+
ind = sampling_logits.topk(topk, dim = dim).indices
|
|
134
|
+
else:
|
|
135
|
+
ind = sampling_logits.argmax(dim = dim)
|
|
136
|
+
|
|
133
137
|
one_hot = F.one_hot(ind, size).type(dtype)
|
|
134
138
|
|
|
135
139
|
if not straight_through or temperature <= 0. or not training:
|
|
@@ -182,7 +186,7 @@ def sample_multinomial(total_count, probs):
|
|
|
182
186
|
return sample.to(device)
|
|
183
187
|
|
|
184
188
|
def all_gather_sizes(x, dim):
|
|
185
|
-
size =
|
|
189
|
+
size = tensor(x.shape[dim], dtype = torch.long, device = x.device)
|
|
186
190
|
all_sizes = [torch.empty_like(size) for _ in range(distributed.get_world_size())]
|
|
187
191
|
distributed.all_gather(all_sizes, size)
|
|
188
192
|
return torch.stack(all_sizes)
|
|
@@ -337,7 +341,7 @@ def orthogonal_loss_fn(t):
|
|
|
337
341
|
|
|
338
342
|
# distance types
|
|
339
343
|
|
|
340
|
-
class
|
|
344
|
+
class Codebook(Module):
|
|
341
345
|
def __init__(
|
|
342
346
|
self,
|
|
343
347
|
dim,
|
|
@@ -359,17 +363,23 @@ class EuclideanCodebook(Module):
|
|
|
359
363
|
affine_param = False,
|
|
360
364
|
sync_affine_param = False,
|
|
361
365
|
affine_param_batch_decay = 0.99,
|
|
362
|
-
affine_param_codebook_decay = 0.9
|
|
366
|
+
affine_param_codebook_decay = 0.9,
|
|
367
|
+
use_cosine_sim = False
|
|
363
368
|
):
|
|
364
369
|
super().__init__()
|
|
365
|
-
self.transform_input = identity
|
|
370
|
+
self.transform_input = identity if not use_cosine_sim else l2norm
|
|
366
371
|
|
|
367
372
|
self.decay = decay
|
|
368
373
|
self.ema_update = ema_update
|
|
369
374
|
self.manual_ema_update = manual_ema_update
|
|
370
375
|
|
|
371
|
-
|
|
372
|
-
|
|
376
|
+
if kmeans_init:
|
|
377
|
+
embed = torch.zeros(num_codebooks, codebook_size, dim)
|
|
378
|
+
else:
|
|
379
|
+
embed = uniform_init(num_codebooks, codebook_size, dim)
|
|
380
|
+
|
|
381
|
+
if use_cosine_sim:
|
|
382
|
+
embed = l2norm(embed)
|
|
373
383
|
|
|
374
384
|
self.codebook_size = codebook_size
|
|
375
385
|
self.num_codebooks = num_codebooks
|
|
@@ -392,7 +402,7 @@ class EuclideanCodebook(Module):
|
|
|
392
402
|
self.kmeans_all_reduce_fn = distributed.all_reduce if use_ddp and sync_kmeans else noop
|
|
393
403
|
self.all_reduce_fn = distributed.all_reduce if use_ddp else noop
|
|
394
404
|
|
|
395
|
-
self.register_buffer('initted',
|
|
405
|
+
self.register_buffer('initted', tensor(not kmeans_init))
|
|
396
406
|
self.register_buffer('cluster_size', torch.ones(num_codebooks, codebook_size))
|
|
397
407
|
self.register_buffer('embed_avg', embed.clone())
|
|
398
408
|
|
|
@@ -402,6 +412,8 @@ class EuclideanCodebook(Module):
|
|
|
402
412
|
else:
|
|
403
413
|
self.register_buffer('embed', embed)
|
|
404
414
|
|
|
415
|
+
self.use_cosine_sim = use_cosine_sim
|
|
416
|
+
|
|
405
417
|
# affine related params
|
|
406
418
|
|
|
407
419
|
self.affine_param = affine_param
|
|
@@ -416,9 +428,9 @@ class EuclideanCodebook(Module):
|
|
|
416
428
|
self.register_buffer('batch_mean', None)
|
|
417
429
|
self.register_buffer('batch_variance', None)
|
|
418
430
|
|
|
419
|
-
self.register_buffer('codebook_mean_needs_init',
|
|
431
|
+
self.register_buffer('codebook_mean_needs_init', tensor(True))
|
|
420
432
|
self.register_buffer('codebook_mean', torch.empty(num_codebooks, 1, dim))
|
|
421
|
-
self.register_buffer('codebook_variance_needs_init',
|
|
433
|
+
self.register_buffer('codebook_variance_needs_init', tensor(True))
|
|
422
434
|
self.register_buffer('codebook_variance', torch.empty(num_codebooks, 1, dim))
|
|
423
435
|
|
|
424
436
|
@torch.jit.ignore
|
|
@@ -434,6 +446,7 @@ class EuclideanCodebook(Module):
|
|
|
434
446
|
data,
|
|
435
447
|
self.codebook_size,
|
|
436
448
|
self.kmeans_iters,
|
|
449
|
+
use_cosine_sim = self.use_cosine_sim,
|
|
437
450
|
sample_fn = self.sample_fn,
|
|
438
451
|
all_reduce_fn = self.kmeans_all_reduce_fn
|
|
439
452
|
)
|
|
@@ -443,7 +456,7 @@ class EuclideanCodebook(Module):
|
|
|
443
456
|
self.embed_avg.data.copy_(embed_sum)
|
|
444
457
|
self.cluster_size.data.copy_(cluster_size)
|
|
445
458
|
self.update_ema()
|
|
446
|
-
self.initted.data.copy_(
|
|
459
|
+
self.initted.data.copy_(tensor(True))
|
|
447
460
|
|
|
448
461
|
@torch.jit.ignore
|
|
449
462
|
def update_with_decay(self, buffer_name, new_value, decay):
|
|
@@ -452,7 +465,7 @@ class EuclideanCodebook(Module):
|
|
|
452
465
|
needs_init = getattr(self, buffer_name + "_needs_init", False)
|
|
453
466
|
|
|
454
467
|
if needs_init:
|
|
455
|
-
self.register_buffer(buffer_name + "_needs_init",
|
|
468
|
+
self.register_buffer(buffer_name + "_needs_init", tensor(False))
|
|
456
469
|
|
|
457
470
|
if not exists(old_value) or needs_init:
|
|
458
471
|
self.register_buffer(buffer_name, new_value.detach())
|
|
@@ -495,7 +508,7 @@ class EuclideanCodebook(Module):
|
|
|
495
508
|
|
|
496
509
|
# number of vectors, for denominator
|
|
497
510
|
|
|
498
|
-
num_vectors =
|
|
511
|
+
num_vectors = tensor(num_vectors, device = device, dtype = dtype)
|
|
499
512
|
distributed.all_reduce(num_vectors)
|
|
500
513
|
|
|
501
514
|
# calculate distributed mean
|
|
@@ -515,6 +528,9 @@ class EuclideanCodebook(Module):
|
|
|
515
528
|
self.update_with_decay('batch_variance', batch_variance, self.affine_param_batch_decay)
|
|
516
529
|
|
|
517
530
|
def replace(self, batch_samples, batch_mask):
|
|
531
|
+
if self.use_cosine_sim:
|
|
532
|
+
batch_samples = l2norm(batch_samples)
|
|
533
|
+
|
|
518
534
|
for ind, (samples, mask) in enumerate(zip(batch_samples, batch_mask)):
|
|
519
535
|
sampled = self.replace_sample_fn(rearrange(samples, '... -> 1 ...'), mask.sum().item())
|
|
520
536
|
sampled = rearrange(sampled, '1 ... -> ...')
|
|
@@ -539,247 +555,74 @@ class EuclideanCodebook(Module):
|
|
|
539
555
|
cluster_size = laplace_smoothing(self.cluster_size, self.codebook_size, self.eps) * self.cluster_size.sum(dim = -1, keepdim = True)
|
|
540
556
|
|
|
541
557
|
embed_normalized = self.embed_avg / rearrange(cluster_size, '... -> ... 1')
|
|
558
|
+
|
|
559
|
+
if self.use_cosine_sim:
|
|
560
|
+
embed_normalized = l2norm(embed_normalized)
|
|
561
|
+
|
|
542
562
|
self.embed.data.copy_(embed_normalized)
|
|
543
563
|
|
|
544
|
-
|
|
545
|
-
def forward(
|
|
564
|
+
def update_ema_part(
|
|
546
565
|
self,
|
|
547
|
-
|
|
548
|
-
|
|
566
|
+
flatten,
|
|
567
|
+
embed_onehot,
|
|
549
568
|
mask = None,
|
|
550
|
-
freeze_codebook = False,
|
|
551
|
-
codebook_transform_fn: Callable | None = None,
|
|
552
569
|
ema_update_weight: Tensor | Callable | None = None,
|
|
553
570
|
accum_ema_update = False
|
|
554
571
|
):
|
|
555
|
-
needs_codebook_dim = x.ndim < 4
|
|
556
|
-
sample_codebook_temp = default(sample_codebook_temp, self.sample_codebook_temp)
|
|
557
|
-
|
|
558
|
-
x = x.float()
|
|
559
|
-
|
|
560
|
-
if needs_codebook_dim:
|
|
561
|
-
x = rearrange(x, '... -> 1 ...')
|
|
562
|
-
|
|
563
|
-
dtype = x.dtype
|
|
564
|
-
flatten, unpack_one = pack_one(x, 'h * d')
|
|
565
|
-
|
|
566
|
-
if exists(mask):
|
|
567
|
-
mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
|
|
568
|
-
|
|
569
|
-
self.init_embed_(flatten, mask = mask)
|
|
570
|
-
|
|
571
|
-
if self.affine_param:
|
|
572
|
-
self.update_affine(flatten, self.embed, mask = mask)
|
|
573
|
-
|
|
574
|
-
# get maybe learnable codes
|
|
575
|
-
|
|
576
|
-
embed = self.embed if self.learnable_codebook else self.embed.detach()
|
|
577
|
-
|
|
578
|
-
embed = embed.to(dtype)
|
|
579
|
-
|
|
580
|
-
# affine params
|
|
581
572
|
if self.affine_param:
|
|
582
573
|
codebook_std = self.codebook_variance.clamp(min = 1e-5).sqrt()
|
|
583
574
|
batch_std = self.batch_variance.clamp(min = 1e-5).sqrt()
|
|
584
|
-
|
|
585
|
-
|
|
586
|
-
# handle maybe implicit neural codebook
|
|
587
|
-
# and calculate distance
|
|
575
|
+
flatten = (flatten - self.batch_mean) * (codebook_std / batch_std) + self.codebook_mean
|
|
588
576
|
|
|
589
|
-
if exists(
|
|
590
|
-
|
|
591
|
-
transformed_embed = rearrange(transformed_embed, 'h b n c d -> h (b n) c d')
|
|
592
|
-
broadcastable_input = rearrange(flatten, '... d -> ... 1 d')
|
|
593
|
-
|
|
594
|
-
dist = -F.pairwise_distance(broadcastable_input, transformed_embed)
|
|
595
|
-
else:
|
|
596
|
-
dist = -cdist(flatten, embed)
|
|
597
|
-
|
|
598
|
-
# sample or argmax depending on temperature
|
|
599
|
-
|
|
600
|
-
embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, temperature = sample_codebook_temp, training = self.training)
|
|
601
|
-
|
|
602
|
-
embed_ind = unpack_one(embed_ind, 'h *')
|
|
577
|
+
if exists(mask):
|
|
578
|
+
embed_onehot[~mask] = 0.
|
|
603
579
|
|
|
604
|
-
|
|
605
|
-
|
|
580
|
+
cluster_size = embed_onehot.sum(dim = 1)
|
|
581
|
+
self.all_reduce_fn(cluster_size)
|
|
606
582
|
|
|
607
|
-
|
|
608
|
-
|
|
583
|
+
embed_sum = einsum('h n d, h n c -> h c d', flatten, embed_onehot)
|
|
584
|
+
embed_sum = embed_sum.contiguous()
|
|
585
|
+
self.all_reduce_fn(embed_sum)
|
|
609
586
|
|
|
610
|
-
|
|
611
|
-
|
|
612
|
-
else:
|
|
613
|
-
quantize = einsum('h b n c, h c d -> h b n d', unpacked_onehot, embed)
|
|
587
|
+
if callable(ema_update_weight):
|
|
588
|
+
ema_update_weight = ema_update_weight(embed_sum, cluster_size)
|
|
614
589
|
|
|
590
|
+
if accum_ema_update:
|
|
591
|
+
accum_grad_(self.cluster_size, cluster_size)
|
|
592
|
+
accum_grad_(self.embed_avg, embed_sum)
|
|
615
593
|
else:
|
|
616
|
-
|
|
617
|
-
|
|
618
|
-
|
|
619
|
-
repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n 1 d', d = transformed_embed.shape[-1])
|
|
620
|
-
quantize = transformed_embed.gather(-2, repeated_embed_ind)
|
|
621
|
-
quantize = rearrange(quantize, 'h b n 1 d -> h b n d')
|
|
622
|
-
|
|
623
|
-
else:
|
|
624
|
-
# quantize = einx.get_at('h [c] d, h b n -> h b n d', embed, embed_ind)
|
|
625
|
-
|
|
626
|
-
repeated_embed = repeat(embed, 'h c d -> h b c d', b = embed_ind.shape[1])
|
|
627
|
-
repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
|
|
628
|
-
quantize = repeated_embed.gather(-2, repeated_embed_ind)
|
|
629
|
-
|
|
630
|
-
if self.training and self.ema_update and not freeze_codebook:
|
|
631
|
-
|
|
632
|
-
if self.affine_param:
|
|
633
|
-
flatten = (flatten - self.batch_mean) * (codebook_std / batch_std) + self.codebook_mean
|
|
634
|
-
|
|
635
|
-
if exists(mask):
|
|
636
|
-
embed_onehot[~mask] = 0.
|
|
637
|
-
|
|
638
|
-
cluster_size = embed_onehot.sum(dim = 1)
|
|
639
|
-
self.all_reduce_fn(cluster_size)
|
|
594
|
+
ema_inplace(self.cluster_size, cluster_size, self.decay, ema_update_weight)
|
|
595
|
+
ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
|
|
640
596
|
|
|
641
|
-
|
|
642
|
-
|
|
643
|
-
|
|
597
|
+
if not self.manual_ema_update:
|
|
598
|
+
self.update_ema()
|
|
599
|
+
self.expire_codes_(flatten)
|
|
644
600
|
|
|
645
|
-
|
|
646
|
-
ema_update_weight = ema_update_weight(embed_sum, cluster_size)
|
|
647
|
-
|
|
648
|
-
if accum_ema_update:
|
|
649
|
-
accum_grad_(self.cluster_size, cluster_size)
|
|
650
|
-
accum_grad_(self.embed_avg, embed_sum)
|
|
651
|
-
else:
|
|
652
|
-
ema_inplace(self.cluster_size, cluster_size, self.decay, ema_update_weight)
|
|
653
|
-
ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
|
|
654
|
-
|
|
655
|
-
if not self.manual_ema_update:
|
|
656
|
-
self.update_ema()
|
|
657
|
-
self.expire_codes_(x)
|
|
658
|
-
|
|
659
|
-
if needs_codebook_dim:
|
|
660
|
-
quantize, embed_ind = map(lambda t: rearrange(t, '1 ... -> ...'), (quantize, embed_ind))
|
|
661
|
-
|
|
662
|
-
dist = unpack_one(dist, 'h * d')
|
|
663
|
-
|
|
664
|
-
return quantize, embed_ind, dist
|
|
665
|
-
|
|
666
|
-
class CosineSimCodebook(Module):
|
|
667
|
-
def __init__(
|
|
601
|
+
def update_ema_indices(
|
|
668
602
|
self,
|
|
669
|
-
|
|
670
|
-
|
|
671
|
-
|
|
672
|
-
|
|
673
|
-
|
|
674
|
-
sync_kmeans = True,
|
|
675
|
-
decay = 0.8,
|
|
676
|
-
eps = 1e-5,
|
|
677
|
-
threshold_ema_dead_code = 2,
|
|
678
|
-
reset_cluster_size = None,
|
|
679
|
-
use_ddp = False,
|
|
680
|
-
learnable_codebook = False,
|
|
681
|
-
gumbel_sample = gumbel_sample,
|
|
682
|
-
sample_codebook_temp = 1.,
|
|
683
|
-
ema_update = True,
|
|
684
|
-
manual_ema_update = False
|
|
603
|
+
x,
|
|
604
|
+
embed_ind,
|
|
605
|
+
mask = None,
|
|
606
|
+
ema_update_weight: Tensor | Callable | None = None,
|
|
607
|
+
accum_ema_update = False
|
|
685
608
|
):
|
|
686
|
-
|
|
687
|
-
|
|
688
|
-
|
|
689
|
-
self.ema_update = ema_update
|
|
690
|
-
self.manual_ema_update = manual_ema_update
|
|
691
|
-
|
|
692
|
-
self.decay = decay
|
|
693
|
-
|
|
694
|
-
if not kmeans_init:
|
|
695
|
-
embed = l2norm(uniform_init(num_codebooks, codebook_size, dim))
|
|
696
|
-
else:
|
|
697
|
-
embed = torch.zeros(num_codebooks, codebook_size, dim)
|
|
698
|
-
|
|
699
|
-
self.codebook_size = codebook_size
|
|
700
|
-
self.num_codebooks = num_codebooks
|
|
701
|
-
|
|
702
|
-
self.kmeans_iters = kmeans_iters
|
|
703
|
-
self.eps = eps
|
|
704
|
-
self.threshold_ema_dead_code = threshold_ema_dead_code
|
|
705
|
-
self.reset_cluster_size = default(reset_cluster_size, threshold_ema_dead_code)
|
|
706
|
-
|
|
707
|
-
assert callable(gumbel_sample)
|
|
708
|
-
self.gumbel_sample = gumbel_sample
|
|
709
|
-
self.sample_codebook_temp = sample_codebook_temp
|
|
710
|
-
|
|
711
|
-
self.sample_fn = sample_vectors_distributed if use_ddp and sync_kmeans else batched_sample_vectors
|
|
712
|
-
|
|
713
|
-
self.replace_sample_fn = sample_vectors_distributed if use_ddp and sync_kmeans else batched_sample_vectors
|
|
714
|
-
|
|
715
|
-
self.kmeans_all_reduce_fn = distributed.all_reduce if use_ddp and sync_kmeans else noop
|
|
716
|
-
self.all_reduce_fn = distributed.all_reduce if use_ddp else noop
|
|
717
|
-
|
|
718
|
-
self.register_buffer('initted', torch.Tensor([not kmeans_init]))
|
|
719
|
-
self.register_buffer('cluster_size', torch.ones(num_codebooks, codebook_size))
|
|
720
|
-
self.register_buffer('embed_avg', embed.clone())
|
|
609
|
+
needs_codebook_dim = x.ndim < 4
|
|
610
|
+
x = x.float()
|
|
721
611
|
|
|
722
|
-
|
|
723
|
-
|
|
724
|
-
self.embed = nn.Parameter(embed)
|
|
725
|
-
else:
|
|
726
|
-
self.register_buffer('embed', embed)
|
|
612
|
+
if needs_codebook_dim:
|
|
613
|
+
x = rearrange(x, '... -> 1 ...')
|
|
727
614
|
|
|
728
|
-
|
|
729
|
-
|
|
730
|
-
if self.initted:
|
|
731
|
-
return
|
|
615
|
+
dtype = x.dtype
|
|
616
|
+
flatten, unpack_one = pack_one(x, 'h * d')
|
|
732
617
|
|
|
733
618
|
if exists(mask):
|
|
734
|
-
c =
|
|
735
|
-
data = rearrange(data[mask], '(c n) d -> c n d', c = c)
|
|
736
|
-
|
|
737
|
-
embed, cluster_size = kmeans(
|
|
738
|
-
data,
|
|
739
|
-
self.codebook_size,
|
|
740
|
-
self.kmeans_iters,
|
|
741
|
-
use_cosine_sim = True,
|
|
742
|
-
sample_fn = self.sample_fn,
|
|
743
|
-
all_reduce_fn = self.kmeans_all_reduce_fn
|
|
744
|
-
)
|
|
745
|
-
|
|
746
|
-
embed_sum = embed * rearrange(cluster_size, '... -> ... 1')
|
|
747
|
-
|
|
748
|
-
self.embed_avg.data.copy_(embed_sum)
|
|
749
|
-
self.cluster_size.data.copy_(cluster_size)
|
|
750
|
-
self.update_ema()
|
|
751
|
-
self.initted.data.copy_(torch.Tensor([True]))
|
|
752
|
-
|
|
753
|
-
def replace(self, batch_samples, batch_mask):
|
|
754
|
-
batch_samples = l2norm(batch_samples)
|
|
755
|
-
|
|
756
|
-
for ind, (samples, mask) in enumerate(zip(batch_samples, batch_mask)):
|
|
757
|
-
sampled = self.replace_sample_fn(rearrange(samples, '... -> 1 ...'), mask.sum().item())
|
|
758
|
-
sampled = rearrange(sampled, '1 ... -> ...')
|
|
759
|
-
|
|
760
|
-
self.embed.data[ind][mask] = sampled
|
|
761
|
-
self.embed_avg.data[ind][mask] = sampled * self.reset_cluster_size
|
|
762
|
-
self.cluster_size.data[ind][mask] = self.reset_cluster_size
|
|
763
|
-
|
|
764
|
-
def expire_codes_(self, batch_samples):
|
|
765
|
-
if self.threshold_ema_dead_code == 0:
|
|
766
|
-
return
|
|
767
|
-
|
|
768
|
-
expired_codes = self.cluster_size < self.threshold_ema_dead_code
|
|
769
|
-
|
|
770
|
-
if not torch.any(expired_codes):
|
|
771
|
-
return
|
|
772
|
-
|
|
773
|
-
batch_samples = rearrange(batch_samples, 'h ... d -> h (...) d')
|
|
774
|
-
self.replace(batch_samples, batch_mask = expired_codes)
|
|
619
|
+
mask = repeat(mask, 'b n -> c (b h n)', c = flatten.shape[0], h = flatten.shape[-2] // (mask.shape[0] * mask.shape[1]))
|
|
775
620
|
|
|
776
|
-
|
|
777
|
-
|
|
621
|
+
embed_ind, _ = pack([embed_ind], 'h *')
|
|
622
|
+
embed_ind = embed_ind.masked_fill(embed_ind == -1, 0)
|
|
623
|
+
embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
|
|
778
624
|
|
|
779
|
-
|
|
780
|
-
embed_normalized = l2norm(embed_normalized)
|
|
781
|
-
|
|
782
|
-
self.embed.data.copy_(embed_normalized)
|
|
625
|
+
self.update_ema_part(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
|
|
783
626
|
|
|
784
627
|
@autocast('cuda', enabled = False)
|
|
785
628
|
def forward(
|
|
@@ -789,9 +632,13 @@ class CosineSimCodebook(Module):
|
|
|
789
632
|
mask = None,
|
|
790
633
|
freeze_codebook = False,
|
|
791
634
|
codebook_transform_fn: Callable | None = None,
|
|
792
|
-
ema_update_weight: Tensor | None = None,
|
|
793
|
-
accum_ema_update = False
|
|
635
|
+
ema_update_weight: Tensor | Callable | None = None,
|
|
636
|
+
accum_ema_update = False,
|
|
637
|
+
ema_update = None,
|
|
638
|
+
topk = None
|
|
794
639
|
):
|
|
640
|
+
ema_update = default(ema_update, self.ema_update)
|
|
641
|
+
|
|
795
642
|
needs_codebook_dim = x.ndim < 4
|
|
796
643
|
sample_codebook_temp = default(sample_codebook_temp, self.sample_codebook_temp)
|
|
797
644
|
|
|
@@ -801,7 +648,6 @@ class CosineSimCodebook(Module):
|
|
|
801
648
|
x = rearrange(x, '... -> 1 ...')
|
|
802
649
|
|
|
803
650
|
dtype = x.dtype
|
|
804
|
-
|
|
805
651
|
flatten, unpack_one = pack_one(x, 'h * d')
|
|
806
652
|
|
|
807
653
|
if exists(mask):
|
|
@@ -809,35 +655,63 @@ class CosineSimCodebook(Module):
|
|
|
809
655
|
|
|
810
656
|
self.init_embed_(flatten, mask = mask)
|
|
811
657
|
|
|
658
|
+
if self.affine_param:
|
|
659
|
+
self.update_affine(flatten, self.embed, mask = mask)
|
|
660
|
+
|
|
661
|
+
# get maybe learnable codes
|
|
662
|
+
|
|
812
663
|
embed = self.embed if self.learnable_codebook else self.embed.detach()
|
|
813
664
|
|
|
814
665
|
embed = embed.to(dtype)
|
|
815
666
|
|
|
667
|
+
# affine params
|
|
668
|
+
|
|
669
|
+
if self.affine_param:
|
|
670
|
+
codebook_std = self.codebook_variance.clamp(min = 1e-5).sqrt()
|
|
671
|
+
batch_std = self.batch_variance.clamp(min = 1e-5).sqrt()
|
|
672
|
+
embed = (embed - self.codebook_mean) * (batch_std / codebook_std) + self.batch_mean
|
|
673
|
+
|
|
816
674
|
# handle maybe implicit neural codebook
|
|
817
|
-
# and
|
|
675
|
+
# and calculate distance
|
|
818
676
|
|
|
819
677
|
if exists(codebook_transform_fn):
|
|
820
678
|
transformed_embed = codebook_transform_fn(embed)
|
|
821
679
|
transformed_embed = rearrange(transformed_embed, 'h b n c d -> h (b n) c d')
|
|
822
|
-
transformed_embed = l2norm(transformed_embed)
|
|
823
680
|
|
|
824
|
-
|
|
681
|
+
if self.use_cosine_sim:
|
|
682
|
+
transformed_embed = l2norm(transformed_embed)
|
|
683
|
+
dist = einsum('h n d, h n c d -> h n c', flatten, transformed_embed)
|
|
684
|
+
else:
|
|
685
|
+
broadcastable_input = rearrange(flatten, '... d -> ... 1 d')
|
|
686
|
+
dist = -F.pairwise_distance(broadcastable_input, transformed_embed)
|
|
825
687
|
else:
|
|
826
|
-
|
|
688
|
+
if self.use_cosine_sim:
|
|
689
|
+
dist = einsum('h n d, h c d -> h n c', flatten, embed)
|
|
690
|
+
else:
|
|
691
|
+
dist = -cdist(flatten, embed)
|
|
692
|
+
|
|
693
|
+
# sample or argmax depending on temperature
|
|
827
694
|
|
|
828
|
-
embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, temperature = sample_codebook_temp, training = self.training)
|
|
829
|
-
|
|
695
|
+
embed_ind, embed_onehot = self.gumbel_sample(dist, dim = -1, topk = topk, temperature = sample_codebook_temp, training = self.training)
|
|
696
|
+
|
|
697
|
+
if exists(topk):
|
|
698
|
+
embed_ind = unpack_one(embed_ind, 'h * k')
|
|
699
|
+
else:
|
|
700
|
+
embed_ind = unpack_one(embed_ind, 'h *')
|
|
830
701
|
|
|
831
702
|
if exists(codebook_transform_fn):
|
|
832
703
|
transformed_embed = unpack_one(transformed_embed, 'h * c d')
|
|
833
704
|
|
|
834
705
|
if self.training:
|
|
835
|
-
|
|
706
|
+
if exists(topk):
|
|
707
|
+
unpacked_onehot = unpack_one(embed_onehot, 'h * k c')
|
|
708
|
+
else:
|
|
709
|
+
unpacked_onehot = unpack_one(embed_onehot, 'h * c')
|
|
836
710
|
|
|
837
711
|
if exists(codebook_transform_fn):
|
|
838
|
-
quantize = einsum('h b n c, h b n c d -> h b n d', unpacked_onehot, transformed_embed)
|
|
712
|
+
quantize = einsum('h b n ... c, h b n c d -> h b n ... d', unpacked_onehot, transformed_embed)
|
|
839
713
|
else:
|
|
840
|
-
quantize = einsum('h b n c, h c d -> h b n d', unpacked_onehot, embed)
|
|
714
|
+
quantize = einsum('h b n ... c, h c d -> h b n ... d', unpacked_onehot, embed)
|
|
841
715
|
|
|
842
716
|
else:
|
|
843
717
|
if exists(codebook_transform_fn):
|
|
@@ -854,36 +728,14 @@ class CosineSimCodebook(Module):
|
|
|
854
728
|
repeated_embed_ind = repeat(embed_ind, 'h b n -> h b n d', d = embed.shape[-1])
|
|
855
729
|
quantize = repeated_embed.gather(-2, repeated_embed_ind)
|
|
856
730
|
|
|
857
|
-
if self.training and
|
|
858
|
-
|
|
859
|
-
embed_onehot[~mask] = 0.
|
|
860
|
-
|
|
861
|
-
bins = embed_onehot.sum(dim = 1)
|
|
862
|
-
self.all_reduce_fn(bins)
|
|
863
|
-
|
|
864
|
-
embed_sum = einsum('h n d, h n c -> h c d', flatten, embed_onehot)
|
|
865
|
-
embed_sum = embed_sum.contiguous()
|
|
866
|
-
self.all_reduce_fn(embed_sum)
|
|
867
|
-
|
|
868
|
-
if callable(ema_update_weight):
|
|
869
|
-
ema_update_weight = ema_update_weight(embed_sum, bins)
|
|
870
|
-
|
|
871
|
-
if accum_ema_update:
|
|
872
|
-
accum_grad_(self.cluster_size, bins)
|
|
873
|
-
accum_grad_(self.embed_avg, embed_sum)
|
|
874
|
-
else:
|
|
875
|
-
|
|
876
|
-
ema_inplace(self.cluster_size, bins, self.decay, ema_update_weight)
|
|
877
|
-
ema_inplace(self.embed_avg, embed_sum, self.decay, ema_update_weight)
|
|
878
|
-
|
|
879
|
-
if not self.manual_ema_update:
|
|
880
|
-
self.update_ema()
|
|
881
|
-
self.expire_codes_(x)
|
|
731
|
+
if self.training and ema_update and not freeze_codebook and not exists(topk):
|
|
732
|
+
self.update_ema_part(flatten, embed_onehot, mask = mask, ema_update_weight = ema_update_weight, accum_ema_update = accum_ema_update)
|
|
882
733
|
|
|
883
734
|
if needs_codebook_dim:
|
|
884
735
|
quantize, embed_ind = map(lambda t: rearrange(t, '1 ... -> ...'), (quantize, embed_ind))
|
|
885
736
|
|
|
886
737
|
dist = unpack_one(dist, 'h * d')
|
|
738
|
+
|
|
887
739
|
return quantize, embed_ind, dist
|
|
888
740
|
|
|
889
741
|
# main class
|
|
@@ -1001,7 +853,7 @@ class VectorQuantize(Module):
|
|
|
1001
853
|
|
|
1002
854
|
self.sync_update_v = sync_update_v
|
|
1003
855
|
|
|
1004
|
-
codebook_class =
|
|
856
|
+
codebook_class = Codebook
|
|
1005
857
|
|
|
1006
858
|
gumbel_sample_fn = partial(
|
|
1007
859
|
gumbel_sample,
|
|
@@ -1027,7 +879,8 @@ class VectorQuantize(Module):
|
|
|
1027
879
|
sample_codebook_temp = sample_codebook_temp,
|
|
1028
880
|
gumbel_sample = gumbel_sample_fn,
|
|
1029
881
|
ema_update = ema_update,
|
|
1030
|
-
manual_ema_update = manual_ema_update
|
|
882
|
+
manual_ema_update = manual_ema_update,
|
|
883
|
+
use_cosine_sim = use_cosine_sim
|
|
1031
884
|
)
|
|
1032
885
|
|
|
1033
886
|
if affine_param:
|
|
@@ -1051,7 +904,7 @@ class VectorQuantize(Module):
|
|
|
1051
904
|
self.accept_image_fmap = accept_image_fmap
|
|
1052
905
|
self.channel_last = channel_last
|
|
1053
906
|
|
|
1054
|
-
self.register_buffer('zero',
|
|
907
|
+
self.register_buffer('zero', tensor(0.), persistent = False)
|
|
1055
908
|
|
|
1056
909
|
# for variable lengthed sequences, whether to take care of masking out the padding to 0 (or return the original input)
|
|
1057
910
|
self.return_zeros_for_masked_padding = return_zeros_for_masked_padding
|
|
@@ -1059,6 +912,10 @@ class VectorQuantize(Module):
|
|
|
1059
912
|
# whether to freeze the codebook, can be overridden on forward
|
|
1060
913
|
self.freeze_codebook = freeze_codebook
|
|
1061
914
|
|
|
915
|
+
@property
|
|
916
|
+
def ema_update(self):
|
|
917
|
+
return self._codebook.ema_update
|
|
918
|
+
|
|
1062
919
|
@property
|
|
1063
920
|
def codebook(self):
|
|
1064
921
|
codebook = self._codebook.embed
|
|
@@ -1133,18 +990,47 @@ class VectorQuantize(Module):
|
|
|
1133
990
|
x = self.maybe_split_heads_from_input(x)
|
|
1134
991
|
self._codebook.expire_codes_(x)
|
|
1135
992
|
|
|
993
|
+
def update_ema_indices(self, x, indices, mask = None):
|
|
994
|
+
if self.accept_image_fmap:
|
|
995
|
+
assert not exists(mask)
|
|
996
|
+
height, width = x.shape[-2:]
|
|
997
|
+
x = rearrange(x, 'b c h w -> b (h w) c')
|
|
998
|
+
|
|
999
|
+
if not self.channel_last and not self.accept_image_fmap:
|
|
1000
|
+
x = rearrange(x, 'b d n -> b n d')
|
|
1001
|
+
|
|
1002
|
+
x = self.project_in(x)
|
|
1003
|
+
x = self.maybe_split_heads_from_input(x)
|
|
1004
|
+
x = self._codebook.transform_input(x)
|
|
1005
|
+
|
|
1006
|
+
if self.heads > 1:
|
|
1007
|
+
if self.separate_codebook_per_head:
|
|
1008
|
+
indices = rearrange(indices, 'b n h -> h b n')
|
|
1009
|
+
else:
|
|
1010
|
+
indices = rearrange(indices, 'b n h -> 1 (b h) n')
|
|
1011
|
+
|
|
1012
|
+
if self.accept_image_fmap:
|
|
1013
|
+
indices = rearrange(indices, 'b h w ... -> b (h w) ...')
|
|
1014
|
+
|
|
1015
|
+
if x.ndim == 2: # only one token
|
|
1016
|
+
indices = rearrange(indices, 'b ... -> b 1 ...')
|
|
1017
|
+
|
|
1018
|
+
self._codebook.update_ema_indices(x, indices, mask = mask)
|
|
1019
|
+
|
|
1136
1020
|
def forward(
|
|
1137
1021
|
self,
|
|
1138
1022
|
x,
|
|
1139
1023
|
indices = None,
|
|
1140
1024
|
mask = None,
|
|
1141
1025
|
lens = None,
|
|
1026
|
+
topk = None,
|
|
1142
1027
|
sample_codebook_temp = None,
|
|
1143
1028
|
freeze_codebook = None,
|
|
1144
1029
|
return_loss_breakdown = False,
|
|
1145
1030
|
codebook_transform_fn: Callable | None = None,
|
|
1146
1031
|
ema_update_weight: Tensor | None = None,
|
|
1147
|
-
accum_ema_update = False
|
|
1032
|
+
accum_ema_update = False,
|
|
1033
|
+
ema_update = None
|
|
1148
1034
|
):
|
|
1149
1035
|
orig_input, input_requires_grad = x, x.requires_grad
|
|
1150
1036
|
|
|
@@ -1202,7 +1088,9 @@ class VectorQuantize(Module):
|
|
|
1202
1088
|
freeze_codebook = freeze_codebook,
|
|
1203
1089
|
codebook_transform_fn = codebook_transform_fn,
|
|
1204
1090
|
ema_update_weight = ema_update_weight,
|
|
1205
|
-
accum_ema_update = accum_ema_update
|
|
1091
|
+
accum_ema_update = accum_ema_update,
|
|
1092
|
+
ema_update = ema_update and not exists(topk),
|
|
1093
|
+
topk = topk
|
|
1206
1094
|
)
|
|
1207
1095
|
|
|
1208
1096
|
# quantize
|
|
@@ -1301,7 +1189,7 @@ class VectorQuantize(Module):
|
|
|
1301
1189
|
|
|
1302
1190
|
# aggregate loss
|
|
1303
1191
|
|
|
1304
|
-
loss =
|
|
1192
|
+
loss = tensor(0., device = device, requires_grad = self.training)
|
|
1305
1193
|
|
|
1306
1194
|
if self.training:
|
|
1307
1195
|
# calculate codebook diversity loss (negative of entropy) if needed
|
|
@@ -1326,9 +1214,19 @@ class VectorQuantize(Module):
|
|
|
1326
1214
|
|
|
1327
1215
|
commit_loss = calculate_ce_loss(embed_ind)
|
|
1328
1216
|
else:
|
|
1329
|
-
if exists(
|
|
1217
|
+
if exists(topk):
|
|
1218
|
+
# handle special case when returning topk
|
|
1219
|
+
|
|
1220
|
+
repeated_input = repeat(orig_input, '... d -> ... k d', k = topk)
|
|
1221
|
+
commit_loss = F.mse_loss(commit_quantize, repeated_input, reduction = 'none')
|
|
1222
|
+
commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
|
|
1223
|
+
|
|
1224
|
+
if exists(mask):
|
|
1225
|
+
commit_loss = einx.where('..., ... k, -> ... k', mask, commit_loss, 0.)
|
|
1226
|
+
|
|
1227
|
+
elif exists(mask):
|
|
1330
1228
|
# with variable lengthed sequences
|
|
1331
|
-
commit_loss = F.mse_loss(commit_quantize,
|
|
1229
|
+
commit_loss = F.mse_loss(commit_quantize, orig_input, reduction = 'none')
|
|
1332
1230
|
|
|
1333
1231
|
loss_mask = mask
|
|
1334
1232
|
if is_multiheaded:
|
|
@@ -1391,7 +1289,7 @@ class VectorQuantize(Module):
|
|
|
1391
1289
|
masked_out_value = torch.zeros_like(orig_input)
|
|
1392
1290
|
|
|
1393
1291
|
quantize = einx.where(
|
|
1394
|
-
'b n, b n d, b n d -> b n d',
|
|
1292
|
+
'b n, b n ... d, b n d -> b n ... d',
|
|
1395
1293
|
mask,
|
|
1396
1294
|
quantize,
|
|
1397
1295
|
masked_out_value
|
{vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/.github/workflows/build.yml
RENAMED
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.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.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_fsq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/examples/autoencoder_lfq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.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.26.0 → vector_quantize_pytorch-1.27.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
|
|
File without changes
|
|
File without changes
|
{vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.26.0 → vector_quantize_pytorch-1.27.0}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|