kernel-elastic-autoencoder 3.2.0__tar.gz → 3.2.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.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 3.2.0
3
+ Version: 3.2.1
4
4
  Summary: Implementation of Kernel-Elastic Autoencoder for Molecular Design (https://doi.org/10.1093/pnasnexus/pgae168)
5
5
  License: MIT
6
6
  Author: Felix Rotter-McCartney
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "kernel_elastic_autoencoder"
3
- version = "3.2.0"
3
+ version = "3.2.1"
4
4
  description = "Implementation of Kernel-Elastic Autoencoder for Molecular Design (https://doi.org/10.1093/pnasnexus/pgae168)"
5
5
  authors = [
6
6
  { name = "Felix Rotter-McCartney", email = "felix.rotter@mail.utoronto.ca" }
@@ -9,6 +9,7 @@ from kernel_elastic_autoencoder.config import TrainingConfig
9
9
  from kernel_elastic_autoencoder.losses import Loss
10
10
  from kernel_elastic_autoencoder.model import Model
11
11
  from kernel_elastic_autoencoder.tokenizer import Tokenizer
12
+ from accelerate.utils import DistributedDataParallelKwargs
12
13
 
13
14
 
14
15
  class Trainer:
@@ -50,7 +51,9 @@ class Trainer:
50
51
  train_split: Fraction of dataset used for training. Must be between 0 and 1.
51
52
  checkpoint: Path of local checkpoint to be saved and/or resumed.
52
53
  """
53
- accelerator = Accelerator()
54
+ accelerator = Accelerator(
55
+ kwargs_handlers=[DistributedDataParallelKwargs(find_unused_parameters=True)]
56
+ )
54
57
 
55
58
  loss_fn = Loss(
56
59
  hp_lambda=self.config_typed.hyperparameters.hp_lambda,
@@ -128,7 +131,7 @@ class Trainer:
128
131
  accelerator.backward(loss)
129
132
  optimizer.step()
130
133
  print(f"Train loss: {float(loss.detach())}")
131
-
134
+
132
135
  model.eval()
133
136
  for input_ids, conditions, token_mask, condition_mask in tqdm(
134
137
  dataloader_test, desc=f"Epoch {epoch}, Test Batch"