vector-quantize-pytorch 1.7.0__tar.gz → 1.7.1__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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: vector_quantize_pytorch
3
- Version: 1.7.0
3
+ Version: 1.7.1
4
4
  Summary: Vector Quantization - Pytorch
5
5
  Home-page: https://github.com/lucidrains/vector-quantizer-pytorch
6
6
  Author: Phil Wang
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
3
3
  setup(
4
4
  name = 'vector_quantize_pytorch',
5
5
  packages = find_packages(),
6
- version = '1.7.0',
6
+ version = '1.7.1',
7
7
  license='MIT',
8
8
  description = 'Vector Quantization - Pytorch',
9
9
  long_description_content_type = 'text/markdown',
@@ -710,7 +710,7 @@ class VectorQuantize(nn.Module):
710
710
  sample_codebook_temp = 1.,
711
711
  straight_through = False,
712
712
  reinmax = False, # using reinmax for improved straight-through, assuming straight through helps at all
713
- sync_codebook = False,
713
+ sync_codebook = None,
714
714
  sync_affine_param = False,
715
715
  ema_update = True,
716
716
  learnable_codebook = False,
@@ -760,6 +760,9 @@ class VectorQuantize(nn.Module):
760
760
  straight_through = straight_through
761
761
  )
762
762
 
763
+ if not exists(sync_codebook):
764
+ sync_codebook = distributed.is_initialized() and distributed.get_world_size() > 1
765
+
763
766
  codebook_kwargs = dict(
764
767
  dim = codebook_dim,
765
768
  num_codebooks = heads if separate_codebook_per_head else 1,
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: vector-quantize-pytorch
3
- Version: 1.7.0
3
+ Version: 1.7.1
4
4
  Summary: Vector Quantization - Pytorch
5
5
  Home-page: https://github.com/lucidrains/vector-quantizer-pytorch
6
6
  Author: Phil Wang