vector-quantize-pytorch 1.31.2__tar.gz → 1.31.4__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.2 → vector_quantize_pytorch-1.31.4}/PKG-INFO +1 -1
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/pyproject.toml +1 -1
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_vq.py +22 -1
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/vector_quantize_pytorch.py +21 -17
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/.gitignore +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/LICENSE +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/README.md +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/__init__.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/binary_mapper.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/evo_lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/finite_scalar_perturbation.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/hierarchical_vq.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/latent_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_fsq.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_lfq.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/sim_vq.py +0 -0
- {vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: vector-quantize-pytorch
|
|
3
|
-
Version: 1.31.
|
|
3
|
+
Version: 1.31.4
|
|
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
|
|
@@ -185,6 +185,7 @@ class ResidualVQ(Module):
|
|
|
185
185
|
eval_beam_size = None,
|
|
186
186
|
beam_score_quantizer_weights: list[float] | None = None,
|
|
187
187
|
quant_grad_frac = 0.,
|
|
188
|
+
return_zeros_for_masked_padding = True,
|
|
188
189
|
**vq_kwargs
|
|
189
190
|
):
|
|
190
191
|
super().__init__()
|
|
@@ -202,6 +203,9 @@ class ResidualVQ(Module):
|
|
|
202
203
|
|
|
203
204
|
self.accept_image_fmap = accept_image_fmap
|
|
204
205
|
|
|
206
|
+
self.return_zeros_for_masked_padding = return_zeros_for_masked_padding
|
|
207
|
+
vq_kwargs.update(return_zeros_for_masked_padding = return_zeros_for_masked_padding)
|
|
208
|
+
|
|
205
209
|
self.implicit_neural_codebook = implicit_neural_codebook
|
|
206
210
|
|
|
207
211
|
if implicit_neural_codebook:
|
|
@@ -397,6 +401,8 @@ class ResidualVQ(Module):
|
|
|
397
401
|
|
|
398
402
|
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
|
|
399
403
|
|
|
404
|
+
orig_input = x
|
|
405
|
+
|
|
400
406
|
beam_size = default(beam_size, self.beam_size if self.training else self.eval_beam_size)
|
|
401
407
|
|
|
402
408
|
is_beam_search = exists(beam_size) and beam_size > 1
|
|
@@ -577,7 +583,7 @@ class ResidualVQ(Module):
|
|
|
577
583
|
|
|
578
584
|
if exists(mask):
|
|
579
585
|
all_losses = einx.where('..., ... l,', mask, all_losses, 0.)
|
|
580
|
-
all_losses = reduce(all_losses, '... l -> l', 'sum') / mask.sum(
|
|
586
|
+
all_losses = reduce(all_losses, '... l -> l', 'sum') / mask.sum().clamp_min(1e-4)
|
|
581
587
|
else:
|
|
582
588
|
all_losses = reduce(all_losses, '... l -> l', 'mean')
|
|
583
589
|
|
|
@@ -609,6 +615,21 @@ class ResidualVQ(Module):
|
|
|
609
615
|
|
|
610
616
|
quantized_out = self.project_out(quantized_out)
|
|
611
617
|
|
|
618
|
+
# if masking, only return quantized for where mask has True
|
|
619
|
+
|
|
620
|
+
if exists(mask):
|
|
621
|
+
masked_out_value = orig_input
|
|
622
|
+
|
|
623
|
+
if self.return_zeros_for_masked_padding:
|
|
624
|
+
masked_out_value = torch.zeros_like(quantized_out)
|
|
625
|
+
|
|
626
|
+
quantized_out = einx.where(
|
|
627
|
+
'b n, b n d, b n d -> b n d',
|
|
628
|
+
mask,
|
|
629
|
+
quantized_out,
|
|
630
|
+
masked_out_value
|
|
631
|
+
)
|
|
632
|
+
|
|
612
633
|
# whether to early return the cross entropy loss
|
|
613
634
|
|
|
614
635
|
if return_loss:
|
|
@@ -1154,6 +1154,7 @@ class VectorQuantize(Module):
|
|
|
1154
1154
|
|
|
1155
1155
|
if need_transpose:
|
|
1156
1156
|
x = rearrange(x, 'b d n -> b n d')
|
|
1157
|
+
orig_input = rearrange(orig_input, 'b d n -> b n d')
|
|
1157
1158
|
|
|
1158
1159
|
# project input
|
|
1159
1160
|
|
|
@@ -1295,7 +1296,7 @@ class VectorQuantize(Module):
|
|
|
1295
1296
|
commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
|
|
1296
1297
|
|
|
1297
1298
|
if exists(mask):
|
|
1298
|
-
commit_loss = einx.where('
|
|
1299
|
+
commit_loss = einx.where('b n, b n ... k, -> b n ... k', mask, commit_loss, 0.)
|
|
1299
1300
|
|
|
1300
1301
|
loss = commit_loss * self.commitment_weight if self.has_commitment_loss else commit_loss
|
|
1301
1302
|
|
|
@@ -1327,7 +1328,7 @@ class VectorQuantize(Module):
|
|
|
1327
1328
|
commit_loss = reduce(commit_loss, '... k d -> ... k', 'mean')
|
|
1328
1329
|
|
|
1329
1330
|
if exists(mask):
|
|
1330
|
-
commit_loss = einx.where('
|
|
1331
|
+
commit_loss = einx.where('b n, b n ... k, -> b n ... k', mask, commit_loss, 0.)
|
|
1331
1332
|
|
|
1332
1333
|
elif exists(mask):
|
|
1333
1334
|
# with variable lengthed sequences
|
|
@@ -1386,20 +1387,6 @@ class VectorQuantize(Module):
|
|
|
1386
1387
|
|
|
1387
1388
|
quantize = self.project_out(quantize)
|
|
1388
1389
|
|
|
1389
|
-
# rearrange quantized embeddings
|
|
1390
|
-
|
|
1391
|
-
if need_transpose:
|
|
1392
|
-
quantize = rearrange(quantize, 'b n ... d -> b d n ...')
|
|
1393
|
-
|
|
1394
|
-
if self.accept_image_fmap:
|
|
1395
|
-
quantize = rearrange(quantize, 'b (h w) ... c -> b c h w ...', h = height, w = width)
|
|
1396
|
-
|
|
1397
|
-
if self.accept_3d_fmap:
|
|
1398
|
-
quantize = rearrange(quantize, "b (d h w) ... c -> b c d h w ...", d=depth, h=height, w=width)
|
|
1399
|
-
|
|
1400
|
-
if only_one:
|
|
1401
|
-
quantize = rearrange(quantize, 'b 1 ... d -> b ... d')
|
|
1402
|
-
|
|
1403
1390
|
# if masking, only return quantized for where mask has True
|
|
1404
1391
|
|
|
1405
1392
|
if exists(mask):
|
|
@@ -1408,8 +1395,11 @@ class VectorQuantize(Module):
|
|
|
1408
1395
|
if self.return_zeros_for_masked_padding:
|
|
1409
1396
|
masked_out_value = torch.zeros_like(orig_input)
|
|
1410
1397
|
|
|
1398
|
+
if is_topk:
|
|
1399
|
+
masked_out_value = repeat(masked_out_value, '... d -> ... k d', k = topk)
|
|
1400
|
+
|
|
1411
1401
|
quantize = einx.where(
|
|
1412
|
-
'b n, b n ... d, b n d -> b n ... d',
|
|
1402
|
+
'b n, b n ... d, b n ... d -> b n ... d',
|
|
1413
1403
|
mask,
|
|
1414
1404
|
quantize,
|
|
1415
1405
|
masked_out_value
|
|
@@ -1422,6 +1412,20 @@ class VectorQuantize(Module):
|
|
|
1422
1412
|
-1
|
|
1423
1413
|
)
|
|
1424
1414
|
|
|
1415
|
+
# rearrange quantized embeddings
|
|
1416
|
+
|
|
1417
|
+
if need_transpose:
|
|
1418
|
+
quantize = rearrange(quantize, 'b n ... d -> b d n ...')
|
|
1419
|
+
|
|
1420
|
+
if self.accept_image_fmap:
|
|
1421
|
+
quantize = rearrange(quantize, 'b (h w) ... c -> b c h w ...', h = height, w = width)
|
|
1422
|
+
|
|
1423
|
+
if self.accept_3d_fmap:
|
|
1424
|
+
quantize = rearrange(quantize, "b (d h w) ... c -> b c d h w ...", d=depth, h=height, w=width)
|
|
1425
|
+
|
|
1426
|
+
if only_one:
|
|
1427
|
+
quantize = rearrange(quantize, 'b 1 ... d -> b ... d')
|
|
1428
|
+
|
|
1425
1429
|
if not return_loss_breakdown:
|
|
1426
1430
|
return quantize, embed_ind, loss
|
|
1427
1431
|
|
|
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
|
|
File without changes
|
{vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/sim_vq.py
RENAMED
|
File without changes
|
{vector_quantize_pytorch-1.31.2 → vector_quantize_pytorch-1.31.4}/vector_quantize_pytorch/utils.py
RENAMED
|
File without changes
|