kernel-elastic-autoencoder 3.2.4__tar.gz → 3.3.0__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.3.0
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.3.0"
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,14 +104,14 @@ 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, loss_fn = (
107
108
  accelerator.prepare(
108
109
  model,
109
110
  optimizer,
110
111
  dataloader_train,
111
112
  dataloader_test,
112
113
  scheduler,
113
- curr_epoch,
114
+ loss_fn,
114
115
  )
115
116
  )
116
117
  accelerator.register_for_checkpointing(scheduler)
@@ -122,6 +123,7 @@ class Trainer:
122
123
  accelerator.wait_for_everyone()
123
124
  for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
124
125
  model.train()
126
+ train_loss = torch.tensor([])
125
127
  for input_ids, conditions, token_mask, condition_mask in tqdm(
126
128
  dataloader_train, desc=f"Epoch {epoch}, Train Batch"
127
129
  ):
@@ -134,9 +136,11 @@ class Trainer:
134
136
  )
135
137
  accelerator.backward(loss)
136
138
  optimizer.step()
137
- accelerator.print(f"Train loss: {float(loss.detach())}")
139
+ torch.cat([train_loss, loss.detach()], dim=0)
140
+ accelerator.print(f"Avg. train loss: {train_loss.mean().item()}")
138
141
 
139
142
  model.eval()
143
+ test_loss = torch.tensor([])
140
144
  for input_ids, conditions, token_mask, condition_mask in tqdm(
141
145
  dataloader_test, desc=f"Epoch {epoch}, Test Batch"
142
146
  ):
@@ -150,7 +154,8 @@ class Trainer:
150
154
  loss = loss_fn(
151
155
  prediction, prediction_noise, input_ids[:, 1:], latents_noise
152
156
  )
153
- accelerator.print(f"Test loss: {float(loss.detach())}")
157
+ torch.cat([test_loss, loss.detach()], dim=0)
158
+ accelerator.print(f"Avg. test loss: {test_loss.mean().item()}")
154
159
 
155
160
  scheduler.step()
156
161