kernel-elastic-autoencoder 3.2.3__tar.gz → 3.2.4__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.3 → kernel_elastic_autoencoder-3.2.4}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/training.py +10 -14
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/README.md +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/losses.py +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/src/kernel_elastic_autoencoder/pipeline.py +0 -0
- {kernel_elastic_autoencoder-3.2.3 → kernel_elastic_autoencoder-3.2.4}/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.4"
|
|
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" }
|
|
@@ -92,10 +92,14 @@ class Trainer:
|
|
|
92
92
|
dataloader_train = torch.utils.data.DataLoader(
|
|
93
93
|
dataset_train,
|
|
94
94
|
batch_size=self.config_typed.common.batch_size,
|
|
95
|
+
pin_memory=True,
|
|
96
|
+
num_workers=4,
|
|
95
97
|
)
|
|
96
98
|
dataloader_test = torch.utils.data.DataLoader(
|
|
97
99
|
dataset_test,
|
|
98
100
|
batch_size=self.config_typed.common.batch_size,
|
|
101
|
+
pin_memory=True,
|
|
102
|
+
num_workers=4,
|
|
99
103
|
)
|
|
100
104
|
curr_epoch = 0
|
|
101
105
|
|
|
@@ -118,11 +122,9 @@ class Trainer:
|
|
|
118
122
|
accelerator.wait_for_everyone()
|
|
119
123
|
for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
|
|
120
124
|
model.train()
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
124
|
-
)
|
|
125
|
-
for input_ids, conditions, token_mask, condition_mask in dataloader_train:
|
|
125
|
+
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
126
|
+
dataloader_train, desc=f"Epoch {epoch}, Train Batch"
|
|
127
|
+
):
|
|
126
128
|
optimizer.zero_grad()
|
|
127
129
|
prediction, prediction_noise, latents_noise = model(
|
|
128
130
|
input_ids, conditions, token_mask, condition_mask
|
|
@@ -132,16 +134,12 @@ class Trainer:
|
|
|
132
134
|
)
|
|
133
135
|
accelerator.backward(loss)
|
|
134
136
|
optimizer.step()
|
|
135
|
-
if accelerator.is_local_main_process:
|
|
136
|
-
batch_bar.update(1)
|
|
137
137
|
accelerator.print(f"Train loss: {float(loss.detach())}")
|
|
138
138
|
|
|
139
139
|
model.eval()
|
|
140
|
-
|
|
141
|
-
|
|
142
|
-
|
|
143
|
-
)
|
|
144
|
-
for input_ids, conditions, token_mask, condition_mask in dataloader_test:
|
|
140
|
+
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
141
|
+
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
142
|
+
):
|
|
145
143
|
with torch.no_grad():
|
|
146
144
|
prediction, prediction_noise, latents_noise = model(
|
|
147
145
|
input_ids,
|
|
@@ -152,8 +150,6 @@ class Trainer:
|
|
|
152
150
|
loss = loss_fn(
|
|
153
151
|
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
154
152
|
)
|
|
155
|
-
if accelerator.is_local_main_process:
|
|
156
|
-
batch_bar.update(1)
|
|
157
153
|
accelerator.print(f"Test loss: {float(loss.detach())}")
|
|
158
154
|
|
|
159
155
|
scheduler.step()
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|