kernel-elastic-autoencoder 3.2.1__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.1 → kernel_elastic_autoencoder-3.2.2}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/training.py +25 -13
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/README.md +0 -0
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/losses.py +0 -0
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.2.1 → kernel_elastic_autoencoder-3.2.2}/src/kernel_elastic_autoencoder/pipeline.py +0 -0
- {kernel_elastic_autoencoder-3.2.1 → 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,13 +3,12 @@ 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
|
|
10
10
|
from kernel_elastic_autoencoder.model import Model
|
|
11
11
|
from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
12
|
-
from accelerate.utils import DistributedDataParallelKwargs
|
|
13
12
|
|
|
14
13
|
|
|
15
14
|
class Trainer:
|
|
@@ -118,9 +117,11 @@ class Trainer:
|
|
|
118
117
|
|
|
119
118
|
for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
|
|
120
119
|
model.train()
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
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:
|
|
124
125
|
optimizer.zero_grad()
|
|
125
126
|
prediction, prediction_noise, latents_noise = model(
|
|
126
127
|
input_ids, conditions, token_mask, condition_mask
|
|
@@ -130,9 +131,15 @@ class Trainer:
|
|
|
130
131
|
)
|
|
131
132
|
accelerator.backward(loss)
|
|
132
133
|
optimizer.step()
|
|
133
|
-
|
|
134
|
+
if accelerator.is_local_main_process:
|
|
135
|
+
batch_bar.update(1)
|
|
136
|
+
accelerator.print(f"Train loss: {float(loss.detach())}")
|
|
134
137
|
|
|
135
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
|
+
)
|
|
136
143
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
137
144
|
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
138
145
|
):
|
|
@@ -146,15 +153,20 @@ class Trainer:
|
|
|
146
153
|
loss = loss_fn(
|
|
147
154
|
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
148
155
|
)
|
|
149
|
-
|
|
156
|
+
if accelerator.is_local_main_process:
|
|
157
|
+
batch_bar.update(1)
|
|
158
|
+
accelerator.print(f"Test loss: {float(loss.detach())}")
|
|
150
159
|
|
|
151
160
|
scheduler.step()
|
|
152
161
|
|
|
153
|
-
accelerator.wait_for_everyone()
|
|
154
162
|
curr_epoch += 1
|
|
155
|
-
|
|
156
|
-
model.save_pretrained(os.path.join(checkpoint, "dist/"))
|
|
157
|
-
self.config_typed.to_json(
|
|
158
|
-
os.path.join(checkpoint, "dist/train_config.json")
|
|
159
|
-
)
|
|
163
|
+
accelerator.wait_for_everyone()
|
|
160
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
|