vector-quantize-pytorch 1.25.0__tar.gz → 1.25.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.
Files changed (33) hide show
  1. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/PKG-INFO +1 -1
  2. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/pyproject.toml +1 -1
  3. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/binary_mapper.py +20 -3
  4. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/.github/workflows/build.yml +0 -0
  5. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/.github/workflows/python-publish.yml +0 -0
  6. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/.github/workflows/test.yml +0 -0
  7. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/.gitignore +0 -0
  8. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/LICENSE +0 -0
  9. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/README.md +0 -0
  10. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/examples/autoencoder.py +0 -0
  11. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/examples/autoencoder_fsq.py +0 -0
  12. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/examples/autoencoder_lfq.py +0 -0
  13. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/examples/autoencoder_sim_vq.py +0 -0
  14. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/images/fsq.png +0 -0
  15. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/images/lfq.png +0 -0
  16. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/images/simvq.png +0 -0
  17. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/images/vq.png +0 -0
  18. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/ruff.toml +0 -0
  19. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/tests/test_latent_quantization.py +0 -0
  20. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/tests/test_lfq.py +0 -0
  21. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/tests/test_readme.py +0 -0
  22. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/__init__.py +0 -0
  23. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/finite_scalar_quantization.py +0 -0
  24. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/latent_quantization.py +0 -0
  25. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/lookup_free_quantization.py +0 -0
  26. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/random_projection_quantizer.py +0 -0
  27. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/residual_fsq.py +0 -0
  28. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/residual_lfq.py +0 -0
  29. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/residual_sim_vq.py +0 -0
  30. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/residual_vq.py +0 -0
  31. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/sim_vq.py +0 -0
  32. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/utils.py +0 -0
  33. {vector_quantize_pytorch-1.25.0 → vector_quantize_pytorch-1.25.2}/vector_quantize_pytorch/vector_quantize_pytorch.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.25.0
3
+ Version: 1.25.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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "vector-quantize-pytorch"
3
- version = "1.25.0"
3
+ version = "1.25.2"
4
4
  description = "Vector Quantization - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -46,7 +46,8 @@ class BinaryMapper(Module):
46
46
  def __init__(
47
47
  self,
48
48
  bits = 1,
49
- kl_loss_threshold = NAT # 1 bit
49
+ kl_loss_threshold = NAT, # 1 bit
50
+ deterministic_on_eval = False
50
51
  ):
51
52
  super().__init__()
52
53
 
@@ -64,14 +65,21 @@ class BinaryMapper(Module):
64
65
  self.kl_loss_threshold = kl_loss_threshold
65
66
  self.register_buffer('zero', tensor(0.), persistent = False)
66
67
 
68
+ # eval behavior
69
+
70
+ self.deterministic_on_eval = deterministic_on_eval
71
+
67
72
  def forward(
68
73
  self,
69
74
  logits,
70
75
  temperature = 1.,
71
76
  straight_through = None,
72
77
  calc_aux_loss = None,
73
- return_indices = False
78
+ deterministic = None,
79
+ return_indices = False,
74
80
  ):
81
+ deterministic = default(deterministic, self.deterministic_on_eval and not self.training)
82
+
75
83
  straight_through = default(straight_through, self.training)
76
84
  calc_aux_loss = default(calc_aux_loss, self.training)
77
85
 
@@ -87,7 +95,11 @@ class BinaryMapper(Module):
87
95
 
88
96
  # sampling
89
97
 
90
- sampled_bits = (torch.rand_like(logits) <= prob_for_sample).long()
98
+ if not deterministic:
99
+ sampled_bits = prob_for_sample.bernoulli().long()
100
+ else:
101
+ sampled_bits = (prob_for_sample > 0.5).long()
102
+
91
103
  indices = (self.power_two * sampled_bits).sum(dim = -1)
92
104
 
93
105
  one_hot = F.one_hot(indices, self.num_codes).float()
@@ -143,3 +155,8 @@ if __name__ == '__main__':
143
155
  assert sparse_one_hot.shape == (3, 4, 2 ** 8)
144
156
  assert indices.shape == (3, 4)
145
157
  assert aux_loss.numel() == 1
158
+
159
+ binary_mapper.eval()
160
+ sparse_one_hot1, _ = binary_mapper(logits, deterministic = True)
161
+ sparse_one_hot2, _ = binary_mapper(logits, deterministic = True)
162
+ assert torch.allclose(sparse_one_hot1, sparse_one_hot2)