kernel-elastic-autoencoder 3.2.4__tar.gz → 3.2.5__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.4
3
+ Version: 3.2.5
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.4"
3
+ version = "3.2.5"
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" }
@@ -197,7 +197,7 @@ class TrainingHyperparameterConfig(Config):
197
197
  r"vanilla-AE and VAE objectives in the reconstruction loss.",
198
198
  )
199
199
  hp_sigma: float = Field(
200
- default=math.sqrt(32),
200
+ default=math.sqrt(0.32),
201
201
  description=r"Hyperparameter $\sigma$, as used in the Kernel function applied in m-MMD loss. Roughly, "
202
202
  r"used as a scaling factor to control the sizes of gradients produced by the m-MMD loss.",
203
203
  )
@@ -86,8 +86,9 @@ class Trainer:
86
86
  dataset = torch.utils.data.TensorDataset(
87
87
  input_ids, conditions, token_mask, condition_mask
88
88
  )
89
+ gen = torch.Generator().manual_seed(0)
89
90
  dataset_train, dataset_test = torch.utils.data.random_split(
90
- dataset, [train_split, 1 - train_split]
91
+ dataset, [train_split, 1 - train_split], generator=gen
91
92
  )
92
93
  dataloader_train = torch.utils.data.DataLoader(
93
94
  dataset_train,
@@ -103,7 +104,7 @@ class Trainer:
103
104
  )
104
105
  curr_epoch = 0
105
106
 
106
- model, optimizer, dataloader_train, dataloader_test, scheduler, curr_epoch = (
107
+ model, optimizer, dataloader_train, dataloader_test, scheduler, curr_epoch, loss_fn = (
107
108
  accelerator.prepare(
108
109
  model,
109
110
  optimizer,
@@ -111,6 +112,7 @@ class Trainer:
111
112
  dataloader_test,
112
113
  scheduler,
113
114
  curr_epoch,
115
+ loss_fn,
114
116
  )
115
117
  )
116
118
  accelerator.register_for_checkpointing(scheduler)