kernel-elastic-autoencoder 3.2.0__tar.gz → 3.2.2__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.
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/training.py +29 -14
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/README.md +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/losses.py +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/pipeline.py +0 -0
- {kernel_elastic_autoencoder-3.2.0 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/tokenizer.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernel_elastic_autoencoder"
|
|
3
|
-
version = "3.2.
|
|
3
|
+
version = "3.2.2"
|
|
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" }
|
|
@@ -3,7 +3,7 @@ from collections.abc import Iterable
|
|
|
3
3
|
|
|
4
4
|
import torch
|
|
5
5
|
from accelerate import Accelerator
|
|
6
|
-
from accelerate.utils import tqdm
|
|
6
|
+
from accelerate.utils import DistributedDataParallelKwargs, tqdm
|
|
7
7
|
|
|
8
8
|
from kernel_elastic_autoencoder.config import TrainingConfig
|
|
9
9
|
from kernel_elastic_autoencoder.losses import Loss
|
|
@@ -50,7 +50,9 @@ class Trainer:
|
|
|
50
50
|
train_split: Fraction of dataset used for training. Must be between 0 and 1.
|
|
51
51
|
checkpoint: Path of local checkpoint to be saved and/or resumed.
|
|
52
52
|
"""
|
|
53
|
-
accelerator = Accelerator(
|
|
53
|
+
accelerator = Accelerator(
|
|
54
|
+
kwargs_handlers=[DistributedDataParallelKwargs(find_unused_parameters=True)]
|
|
55
|
+
)
|
|
54
56
|
|
|
55
57
|
loss_fn = Loss(
|
|
56
58
|
hp_lambda=self.config_typed.hyperparameters.hp_lambda,
|
|
@@ -115,9 +117,11 @@ class Trainer:
|
|
|
115
117
|
|
|
116
118
|
for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
|
|
117
119
|
model.train()
|
|
118
|
-
|
|
119
|
-
|
|
120
|
-
|
|
120
|
+
if accelerator.is_local_main_process:
|
|
121
|
+
batch_bar = tqdm(
|
|
122
|
+
total=len(dataloader_train), desc=f"Epoch {epoch}, Train Batch"
|
|
123
|
+
)
|
|
124
|
+
for input_ids, conditions, token_mask, condition_mask in dataloader_train:
|
|
121
125
|
optimizer.zero_grad()
|
|
122
126
|
prediction, prediction_noise, latents_noise = model(
|
|
123
127
|
input_ids, conditions, token_mask, condition_mask
|
|
@@ -127,9 +131,15 @@ class Trainer:
|
|
|
127
131
|
)
|
|
128
132
|
accelerator.backward(loss)
|
|
129
133
|
optimizer.step()
|
|
130
|
-
|
|
131
|
-
|
|
134
|
+
if accelerator.is_local_main_process:
|
|
135
|
+
batch_bar.update(1)
|
|
136
|
+
accelerator.print(f"Train loss: {float(loss.detach())}")
|
|
137
|
+
|
|
132
138
|
model.eval()
|
|
139
|
+
if accelerator.is_local_main_process:
|
|
140
|
+
batch_bar = tqdm(
|
|
141
|
+
total=len(dataloader_train), desc=f"Epoch {epoch}, Test Batch"
|
|
142
|
+
)
|
|
133
143
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
134
144
|
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
135
145
|
):
|
|
@@ -143,15 +153,20 @@ class Trainer:
|
|
|
143
153
|
loss = loss_fn(
|
|
144
154
|
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
145
155
|
)
|
|
146
|
-
|
|
156
|
+
if accelerator.is_local_main_process:
|
|
157
|
+
batch_bar.update(1)
|
|
158
|
+
accelerator.print(f"Test loss: {float(loss.detach())}")
|
|
147
159
|
|
|
148
160
|
scheduler.step()
|
|
149
161
|
|
|
150
|
-
accelerator.wait_for_everyone()
|
|
151
162
|
curr_epoch += 1
|
|
152
|
-
|
|
153
|
-
model.save_pretrained(os.path.join(checkpoint, "dist/"))
|
|
154
|
-
self.config_typed.to_json(
|
|
155
|
-
os.path.join(checkpoint, "dist/train_config.json")
|
|
156
|
-
)
|
|
163
|
+
accelerator.wait_for_everyone()
|
|
157
164
|
accelerator.save_state(checkpoint)
|
|
165
|
+
if accelerator.is_main_process:
|
|
166
|
+
os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
|
|
167
|
+
accelerator.unwrap_model(model).save_pretrained(
|
|
168
|
+
os.path.join(checkpoint, "dist/")
|
|
169
|
+
)
|
|
170
|
+
self.config_typed.to_json(
|
|
171
|
+
os.path.join(checkpoint, "dist/train_config.json")
|
|
172
|
+
)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|