kernel-elastic-autoencoder 3.3.6__tar.gz → 3.4.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.
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/losses.py +6 -7
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/pipeline.py +6 -4
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/training.py +40 -5
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/README.md +0 -0
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.3.6 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernel_elastic_autoencoder"
|
|
3
|
-
version = "3.
|
|
3
|
+
version = "3.4.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" }
|
|
@@ -61,8 +61,11 @@ class Loss(nn.Module):
|
|
|
61
61
|
"""Sequence length dimension to which inputs are pooled after condition concatenation through the
|
|
62
62
|
encoder. Proportional to the dimension of latent vectors."""
|
|
63
63
|
|
|
64
|
-
|
|
65
|
-
|
|
64
|
+
mvn = torch.distributions.MultivariateNormal(
|
|
65
|
+
loc=torch.zeros(self.pooling_dim * self.embedding_dim),
|
|
66
|
+
covariance_matrix=torch.eye(self.pooling_dim * self.embedding_dim),
|
|
67
|
+
)
|
|
68
|
+
self.register_buffer("_samples", mvn.sample((self.kernel_dist_size,)))
|
|
66
69
|
|
|
67
70
|
def forward(
|
|
68
71
|
self,
|
|
@@ -110,11 +113,7 @@ class Loss(nn.Module):
|
|
|
110
113
|
self,
|
|
111
114
|
latents: torch.Tensor,
|
|
112
115
|
) -> torch.Tensor:
|
|
113
|
-
|
|
114
|
-
loc=self._loc, # type: ignore
|
|
115
|
-
covariance_matrix=self._cov, # type: ignore
|
|
116
|
-
)
|
|
117
|
-
samples = mvn.rsample((self.kernel_dist_size,)).to(latents.device)
|
|
116
|
+
samples = self._samples
|
|
118
117
|
square_difference_sum = torch.cdist(latents, samples, p=2.0).pow(2)
|
|
119
118
|
kernel_pairwise_sum = torch.exp(
|
|
120
119
|
((-1 / (self.pooling_dim * self.embedding_dim)) * square_difference_sum)
|
|
@@ -52,6 +52,7 @@ class Pipeline:
|
|
|
52
52
|
latents: torch.Tensor,
|
|
53
53
|
sequences: list[str],
|
|
54
54
|
conditions: list[list[float]] | torch.Tensor,
|
|
55
|
+
with_grad: bool = False,
|
|
55
56
|
**kwargs,
|
|
56
57
|
) -> Iterable[str]:
|
|
57
58
|
"""Completes each conditioned input sequence.
|
|
@@ -181,14 +182,15 @@ class Pipeline:
|
|
|
181
182
|
condition_embeddings=conds_embed,
|
|
182
183
|
token_mask=token_mask,
|
|
183
184
|
)
|
|
185
|
+
odds = logits.log_softmax(dim=-1)
|
|
184
186
|
new_toks = (
|
|
185
|
-
torch.topk(
|
|
187
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
186
188
|
.indices.flatten()
|
|
187
189
|
.unsqueeze(-1)
|
|
188
190
|
.to(torch.long)
|
|
189
191
|
)
|
|
190
192
|
new_probs = (
|
|
191
|
-
torch.topk(
|
|
193
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
192
194
|
.values.flatten()
|
|
193
195
|
.unsqueeze(-1)
|
|
194
196
|
)
|
|
@@ -263,7 +265,7 @@ class Pipeline:
|
|
|
263
265
|
top_prob_inds = (
|
|
264
266
|
(
|
|
265
267
|
top_probs.sum(dim=-1)
|
|
266
|
-
|
|
268
|
+
/ torch.sqrt(grouped_ids != self.tokenizer.pad_token_id)
|
|
267
269
|
.to(torch.long)
|
|
268
270
|
.sum(dim=-1)
|
|
269
271
|
)
|
|
@@ -295,7 +297,7 @@ class Pipeline:
|
|
|
295
297
|
top_prob_inds = (
|
|
296
298
|
(
|
|
297
299
|
top_probs.sum(dim=-1)
|
|
298
|
-
|
|
300
|
+
/ torch.sqrt(grouped_ids != self.tokenizer.pad_token_id)
|
|
299
301
|
.to(torch.long)
|
|
300
302
|
.sum(dim=-1)
|
|
301
303
|
)
|
|
@@ -1,5 +1,7 @@
|
|
|
1
|
+
import json
|
|
1
2
|
import os
|
|
2
|
-
from collections.abc import Iterable
|
|
3
|
+
from collections.abc import Callable, Iterable
|
|
4
|
+
from typing import Any
|
|
3
5
|
|
|
4
6
|
import torch
|
|
5
7
|
from accelerate import Accelerator
|
|
@@ -37,6 +39,7 @@ class Trainer:
|
|
|
37
39
|
conditions: Iterable[Iterable[float]] | torch.Tensor,
|
|
38
40
|
train_split: float = 0.9,
|
|
39
41
|
checkpoint: str = "./checkpoint",
|
|
42
|
+
epoch_callback: Callable[[dict[str, Any]], Any] = lambda _: None,
|
|
40
43
|
):
|
|
41
44
|
"""Trains a model. Optionally, resumes from an existing checkpoint.
|
|
42
45
|
|
|
@@ -49,6 +52,7 @@ class Trainer:
|
|
|
49
52
|
conditions: Iterable of iterables of condition values per sequence.
|
|
50
53
|
train_split: Fraction of dataset used for training. Must be between 0 and 1.
|
|
51
54
|
checkpoint: Path of local checkpoint to be saved and/or resumed.
|
|
55
|
+
epoch_callback: Callback function accepting a dict of per-epoch stats.
|
|
52
56
|
"""
|
|
53
57
|
accelerator = Accelerator(
|
|
54
58
|
kwargs_handlers=[DistributedDataParallelKwargs(find_unused_parameters=True)]
|
|
@@ -123,7 +127,16 @@ class Trainer:
|
|
|
123
127
|
curr_epoch = scheduler.scheduler.last_epoch + 1
|
|
124
128
|
|
|
125
129
|
accelerator.wait_for_everyone()
|
|
130
|
+
all_contexts = []
|
|
126
131
|
for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
|
|
132
|
+
if accelerator.is_main_process:
|
|
133
|
+
cb_ctx = {
|
|
134
|
+
"epoch": epoch,
|
|
135
|
+
"train_loss": None,
|
|
136
|
+
"test_loss": None,
|
|
137
|
+
"dist_mean": None,
|
|
138
|
+
"dist_var": None,
|
|
139
|
+
}
|
|
127
140
|
model.train()
|
|
128
141
|
train_loss = torch.tensor([], device=accelerator.device)
|
|
129
142
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
@@ -133,16 +146,21 @@ class Trainer:
|
|
|
133
146
|
prediction, prediction_noise, latents = model(
|
|
134
147
|
input_ids, conditions, token_mask, condition_mask
|
|
135
148
|
)
|
|
136
|
-
loss = loss_fn(
|
|
137
|
-
prediction, prediction_noise, input_ids[:, 1:], latents
|
|
138
|
-
)
|
|
149
|
+
loss = loss_fn(prediction, prediction_noise, input_ids[:, 1:], latents)
|
|
139
150
|
accelerator.backward(loss)
|
|
151
|
+
if accelerator.sync_gradients:
|
|
152
|
+
accelerator.clip_grad_norm_(model.parameters(), 1.0)
|
|
140
153
|
optimizer.step()
|
|
141
154
|
train_loss = torch.cat([train_loss, loss.detach().unsqueeze(-1)], dim=0)
|
|
142
155
|
accelerator.print(f"Avg. train loss: {train_loss.mean().item()}")
|
|
156
|
+
if accelerator.is_main_process:
|
|
157
|
+
cb_ctx["train_loss"] = train_loss.mean().item()
|
|
143
158
|
|
|
144
159
|
model.eval()
|
|
145
160
|
test_loss = torch.tensor([], device=accelerator.device)
|
|
161
|
+
|
|
162
|
+
dist_mean = torch.tensor([], device=accelerator.device)
|
|
163
|
+
dist_var = torch.tensor([], device=accelerator.device)
|
|
146
164
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
147
165
|
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
148
166
|
):
|
|
@@ -156,9 +174,22 @@ class Trainer:
|
|
|
156
174
|
loss = loss_fn(
|
|
157
175
|
prediction, prediction_noise, input_ids[:, 1:], latents
|
|
158
176
|
)
|
|
159
|
-
|
|
177
|
+
mean = latents.mean().mean()
|
|
178
|
+
var = latents.var(dim=0).mean()
|
|
179
|
+
test_loss = torch.cat(
|
|
180
|
+
[test_loss, loss.detach().unsqueeze(-1)], dim=0
|
|
181
|
+
)
|
|
182
|
+
dist_mean = torch.cat(
|
|
183
|
+
[dist_mean, mean.detach().unsqueeze(-1)], dim=0
|
|
184
|
+
)
|
|
185
|
+
dist_var = torch.cat([dist_var, var.detach().unsqueeze(-1)], dim=0)
|
|
160
186
|
accelerator.print(f"Avg. test loss: {test_loss.mean().item()}")
|
|
161
187
|
|
|
188
|
+
if accelerator.is_main_process:
|
|
189
|
+
cb_ctx["test_loss"] = test_loss.mean().item()
|
|
190
|
+
cb_ctx["dist_mean"] = dist_mean.mean().item()
|
|
191
|
+
cb_ctx["dist_var"] = dist_var.mean().item()
|
|
192
|
+
|
|
162
193
|
scheduler.step()
|
|
163
194
|
|
|
164
195
|
curr_epoch += 1
|
|
@@ -166,6 +197,10 @@ class Trainer:
|
|
|
166
197
|
accelerator.save_state(checkpoint)
|
|
167
198
|
if accelerator.is_main_process:
|
|
168
199
|
os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
|
|
200
|
+
all_contexts.append(cb_ctx)
|
|
201
|
+
with open("log.json", "w") as f:
|
|
202
|
+
json.dump(all_contexts, f, indent=4)
|
|
203
|
+
epoch_callback(cb_ctx)
|
|
169
204
|
accelerator.unwrap_model(model).save_pretrained(
|
|
170
205
|
os.path.join(checkpoint, "dist/")
|
|
171
206
|
)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|