kernel-elastic-autoencoder 3.2.3__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.3
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.3"
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,20 +86,25 @@ 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,
94
95
  batch_size=self.config_typed.common.batch_size,
96
+ pin_memory=True,
97
+ num_workers=4,
95
98
  )
96
99
  dataloader_test = torch.utils.data.DataLoader(
97
100
  dataset_test,
98
101
  batch_size=self.config_typed.common.batch_size,
102
+ pin_memory=True,
103
+ num_workers=4,
99
104
  )
100
105
  curr_epoch = 0
101
106
 
102
- model, optimizer, dataloader_train, dataloader_test, scheduler, curr_epoch = (
107
+ model, optimizer, dataloader_train, dataloader_test, scheduler, curr_epoch, loss_fn = (
103
108
  accelerator.prepare(
104
109
  model,
105
110
  optimizer,
@@ -107,6 +112,7 @@ class Trainer:
107
112
  dataloader_test,
108
113
  scheduler,
109
114
  curr_epoch,
115
+ loss_fn,
110
116
  )
111
117
  )
112
118
  accelerator.register_for_checkpointing(scheduler)
@@ -118,11 +124,9 @@ class Trainer:
118
124
  accelerator.wait_for_everyone()
119
125
  for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
120
126
  model.train()
121
- if accelerator.is_local_main_process:
122
- batch_bar = tqdm(
123
- total=len(dataloader_train), desc=f"Epoch {epoch}, Train Batch"
124
- )
125
- for input_ids, conditions, token_mask, condition_mask in dataloader_train:
127
+ for input_ids, conditions, token_mask, condition_mask in tqdm(
128
+ dataloader_train, desc=f"Epoch {epoch}, Train Batch"
129
+ ):
126
130
  optimizer.zero_grad()
127
131
  prediction, prediction_noise, latents_noise = model(
128
132
  input_ids, conditions, token_mask, condition_mask
@@ -132,16 +136,12 @@ class Trainer:
132
136
  )
133
137
  accelerator.backward(loss)
134
138
  optimizer.step()
135
- if accelerator.is_local_main_process:
136
- batch_bar.update(1)
137
139
  accelerator.print(f"Train loss: {float(loss.detach())}")
138
140
 
139
141
  model.eval()
140
- if accelerator.is_local_main_process:
141
- batch_bar = tqdm(
142
- total=len(dataloader_train), desc=f"Epoch {epoch}, Test Batch"
143
- )
144
- for input_ids, conditions, token_mask, condition_mask in dataloader_test:
142
+ for input_ids, conditions, token_mask, condition_mask in tqdm(
143
+ dataloader_test, desc=f"Epoch {epoch}, Test Batch"
144
+ ):
145
145
  with torch.no_grad():
146
146
  prediction, prediction_noise, latents_noise = model(
147
147
  input_ids,
@@ -152,8 +152,6 @@ class Trainer:
152
152
  loss = loss_fn(
153
153
  prediction, prediction_noise, input_ids[:, 1:], latents_noise
154
154
  )
155
- if accelerator.is_local_main_process:
156
- batch_bar.update(1)
157
155
  accelerator.print(f"Test loss: {float(loss.detach())}")
158
156
 
159
157
  scheduler.step()