kernel-elastic-autoencoder 3.3.5__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.5 → kernel_elastic_autoencoder-3.4.0}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/losses.py +6 -7
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/pipeline.py +62 -23
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/training.py +44 -7
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/README.md +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.4.0}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → 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.
|
|
@@ -108,7 +109,7 @@ class Pipeline:
|
|
|
108
109
|
.indices.squeeze(-1)
|
|
109
110
|
.to(torch.long)
|
|
110
111
|
)
|
|
111
|
-
batches_completed |= new_toks.
|
|
112
|
+
batches_completed |= new_toks.flatten() == self.tokenizer.eos_token_id
|
|
112
113
|
new_toks = torch.where(
|
|
113
114
|
batches_completed.unsqueeze(-1),
|
|
114
115
|
self.tokenizer.pad_token_id,
|
|
@@ -118,9 +119,7 @@ class Pipeline:
|
|
|
118
119
|
token_mask = torch.cat(
|
|
119
120
|
[
|
|
120
121
|
token_mask,
|
|
121
|
-
|
|
122
|
-
input_ids.size(0), 1, dtype=torch.bool, device=self.device
|
|
123
|
-
),
|
|
122
|
+
batches_completed.unsqueeze(-1),
|
|
124
123
|
],
|
|
125
124
|
dim=1,
|
|
126
125
|
)
|
|
@@ -136,10 +135,9 @@ class Pipeline:
|
|
|
136
135
|
) -> Iterable[str]:
|
|
137
136
|
"""Completes each conditioned input sequence, using the beam search strategy.
|
|
138
137
|
|
|
139
|
-
The model completes each sequence in the provided list using decoder-only inference.
|
|
140
|
-
|
|
141
|
-
|
|
142
|
-
the highest sum of token logits is chosen for each batch.
|
|
138
|
+
The model completes each sequence in the provided list using decoder-only inference. The beam search
|
|
139
|
+
decoding algorithm is used. On the last step, instead of returning beam_size candidates per batch, the
|
|
140
|
+
candidate with the highest length-normalized sum of log odds is selected for each batch.
|
|
143
141
|
|
|
144
142
|
Args:
|
|
145
143
|
latents: Tensor of dimension (B, P * E) containing latent vectors for the batch.
|
|
@@ -184,14 +182,15 @@ class Pipeline:
|
|
|
184
182
|
condition_embeddings=conds_embed,
|
|
185
183
|
token_mask=token_mask,
|
|
186
184
|
)
|
|
185
|
+
odds = logits.log_softmax(dim=-1)
|
|
187
186
|
new_toks = (
|
|
188
|
-
torch.topk(
|
|
187
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
189
188
|
.indices.flatten()
|
|
190
189
|
.unsqueeze(-1)
|
|
191
190
|
.to(torch.long)
|
|
192
191
|
)
|
|
193
192
|
new_probs = (
|
|
194
|
-
torch.topk(
|
|
193
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
195
194
|
.values.flatten()
|
|
196
195
|
.unsqueeze(-1)
|
|
197
196
|
)
|
|
@@ -212,11 +211,11 @@ class Pipeline:
|
|
|
212
211
|
conds_embed = conds_embed.repeat_interleave(beam_size, dim=0)
|
|
213
212
|
latents = latents.repeat_interleave(beam_size, dim=0)
|
|
214
213
|
token_mask = token_mask.repeat_interleave(beam_size, dim=0)
|
|
215
|
-
batches_completed |= new_toks.
|
|
214
|
+
batches_completed |= new_toks.flatten() == self.tokenizer.eos_token_id
|
|
216
215
|
token_mask = torch.cat(
|
|
217
216
|
[
|
|
218
217
|
token_mask,
|
|
219
|
-
|
|
218
|
+
batches_completed.unsqueeze(-1),
|
|
220
219
|
],
|
|
221
220
|
dim=1,
|
|
222
221
|
)
|
|
@@ -230,41 +229,81 @@ class Pipeline:
|
|
|
230
229
|
condition_embeddings=conds_embed,
|
|
231
230
|
token_mask=token_mask,
|
|
232
231
|
)
|
|
232
|
+
odds = logits.log_softmax(dim=-1)
|
|
233
233
|
new_toks = (
|
|
234
|
-
torch.topk(
|
|
234
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
235
235
|
.indices.flatten()
|
|
236
236
|
.unsqueeze(-1)
|
|
237
237
|
.to(torch.long)
|
|
238
238
|
)
|
|
239
239
|
new_probs = (
|
|
240
|
-
torch.topk(
|
|
240
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
241
|
+
.values.flatten()
|
|
242
|
+
.unsqueeze(-1)
|
|
241
243
|
)
|
|
242
244
|
new_toks = torch.where(
|
|
243
|
-
batches_completed.unsqueeze(-1),
|
|
245
|
+
batches_completed.repeat_interleave(beam_size, dim=0).unsqueeze(-1),
|
|
244
246
|
self.tokenizer.pad_token_id,
|
|
245
247
|
new_toks,
|
|
246
248
|
)
|
|
247
249
|
new_probs = torch.where(
|
|
248
|
-
batches_completed.unsqueeze(-1),
|
|
250
|
+
batches_completed.repeat_interleave(beam_size, dim=0).unsqueeze(-1),
|
|
249
251
|
0.0,
|
|
250
252
|
new_probs,
|
|
251
253
|
)
|
|
252
|
-
|
|
253
|
-
|
|
254
|
+
|
|
255
|
+
candidate_ids = torch.cat(
|
|
256
|
+
[input_ids.repeat_interleave(beam_size, dim=0), new_toks], dim=1
|
|
257
|
+
)
|
|
258
|
+
candidate_probs = torch.cat(
|
|
259
|
+
[input_probs.repeat_interleave(beam_size, dim=0), new_probs], dim=1
|
|
260
|
+
)
|
|
261
|
+
top_probs = candidate_probs.view(
|
|
262
|
+
candidate_probs.size(0) // (beam_size**2), beam_size**2, -1
|
|
263
|
+
)
|
|
264
|
+
grouped_ids = candidate_ids.view(top_probs.size(0), top_probs.size(1), -1)
|
|
265
|
+
top_prob_inds = (
|
|
266
|
+
(
|
|
267
|
+
top_probs.sum(dim=-1)
|
|
268
|
+
/ torch.sqrt(grouped_ids != self.tokenizer.pad_token_id)
|
|
269
|
+
.to(torch.long)
|
|
270
|
+
.sum(dim=-1)
|
|
271
|
+
)
|
|
272
|
+
.topk(k=beam_size, dim=1)
|
|
273
|
+
.indices.squeeze(-1)
|
|
274
|
+
)
|
|
275
|
+
input_ids = grouped_ids[
|
|
276
|
+
torch.arange(top_probs.size(0)).unsqueeze(-1).repeat(1, beam_size),
|
|
277
|
+
top_prob_inds,
|
|
278
|
+
].view(input_ids.size(0), -1)
|
|
279
|
+
input_probs = top_probs[
|
|
280
|
+
torch.arange(top_probs.size(0)).unsqueeze(-1).repeat(1, beam_size),
|
|
281
|
+
top_prob_inds,
|
|
282
|
+
].view(input_ids.size(0), -1)
|
|
283
|
+
|
|
284
|
+
batches_completed |= (
|
|
285
|
+
input_ids[:, -1:].flatten() == self.tokenizer.eos_token_id
|
|
286
|
+
)
|
|
254
287
|
token_mask = torch.cat(
|
|
255
288
|
[
|
|
256
289
|
token_mask,
|
|
257
|
-
|
|
258
|
-
input_ids.size(0), 1, dtype=torch.bool, device=self.device
|
|
259
|
-
),
|
|
290
|
+
batches_completed.unsqueeze(-1),
|
|
260
291
|
],
|
|
261
292
|
dim=1,
|
|
262
293
|
)
|
|
263
|
-
batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
264
294
|
|
|
265
295
|
top_probs = input_probs.reshape(input_probs.size(0) // beam_size, beam_size, -1)
|
|
266
|
-
top_prob_inds = top_probs.sum(dim=-1).topk(k=1, dim=1).indices.squeeze(-1)
|
|
267
296
|
grouped_ids = input_ids.view(top_probs.shape[0], beam_size, -1)
|
|
297
|
+
top_prob_inds = (
|
|
298
|
+
(
|
|
299
|
+
top_probs.sum(dim=-1)
|
|
300
|
+
/ torch.sqrt(grouped_ids != self.tokenizer.pad_token_id)
|
|
301
|
+
.to(torch.long)
|
|
302
|
+
.sum(dim=-1)
|
|
303
|
+
)
|
|
304
|
+
.topk(k=1, dim=1)
|
|
305
|
+
.indices.squeeze(-1)
|
|
306
|
+
)
|
|
268
307
|
winning_ids = grouped_ids[torch.arange(top_probs.size(0)), top_prob_inds]
|
|
269
308
|
return self.tokenizer.decode(winning_ids, skip_special_tokens=True)
|
|
270
309
|
|
|
@@ -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)]
|
|
@@ -94,13 +98,15 @@ class Trainer:
|
|
|
94
98
|
dataset_train,
|
|
95
99
|
batch_size=self.config_typed.common.batch_size,
|
|
96
100
|
pin_memory=True,
|
|
97
|
-
num_workers=
|
|
101
|
+
num_workers=4,
|
|
102
|
+
shuffle=True,
|
|
98
103
|
)
|
|
99
104
|
dataloader_test = torch.utils.data.DataLoader(
|
|
100
105
|
dataset_test,
|
|
101
106
|
batch_size=self.config_typed.common.batch_size,
|
|
102
107
|
pin_memory=True,
|
|
103
|
-
num_workers=
|
|
108
|
+
num_workers=4,
|
|
109
|
+
shuffle=True,
|
|
104
110
|
)
|
|
105
111
|
curr_epoch = 0
|
|
106
112
|
|
|
@@ -121,7 +127,16 @@ class Trainer:
|
|
|
121
127
|
curr_epoch = scheduler.scheduler.last_epoch + 1
|
|
122
128
|
|
|
123
129
|
accelerator.wait_for_everyone()
|
|
130
|
+
all_contexts = []
|
|
124
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
|
+
}
|
|
125
140
|
model.train()
|
|
126
141
|
train_loss = torch.tensor([], device=accelerator.device)
|
|
127
142
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
@@ -131,16 +146,21 @@ class Trainer:
|
|
|
131
146
|
prediction, prediction_noise, latents = model(
|
|
132
147
|
input_ids, conditions, token_mask, condition_mask
|
|
133
148
|
)
|
|
134
|
-
loss = loss_fn(
|
|
135
|
-
prediction, prediction_noise, input_ids[:, 1:], latents
|
|
136
|
-
)
|
|
149
|
+
loss = loss_fn(prediction, prediction_noise, input_ids[:, 1:], latents)
|
|
137
150
|
accelerator.backward(loss)
|
|
151
|
+
if accelerator.sync_gradients:
|
|
152
|
+
accelerator.clip_grad_norm_(model.parameters(), 1.0)
|
|
138
153
|
optimizer.step()
|
|
139
154
|
train_loss = torch.cat([train_loss, loss.detach().unsqueeze(-1)], dim=0)
|
|
140
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()
|
|
141
158
|
|
|
142
159
|
model.eval()
|
|
143
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)
|
|
144
164
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
145
165
|
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
146
166
|
):
|
|
@@ -154,9 +174,22 @@ class Trainer:
|
|
|
154
174
|
loss = loss_fn(
|
|
155
175
|
prediction, prediction_noise, input_ids[:, 1:], latents
|
|
156
176
|
)
|
|
157
|
-
|
|
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)
|
|
158
186
|
accelerator.print(f"Avg. test loss: {test_loss.mean().item()}")
|
|
159
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
|
+
|
|
160
193
|
scheduler.step()
|
|
161
194
|
|
|
162
195
|
curr_epoch += 1
|
|
@@ -164,6 +197,10 @@ class Trainer:
|
|
|
164
197
|
accelerator.save_state(checkpoint)
|
|
165
198
|
if accelerator.is_main_process:
|
|
166
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)
|
|
167
204
|
accelerator.unwrap_model(model).save_pretrained(
|
|
168
205
|
os.path.join(checkpoint, "dist/")
|
|
169
206
|
)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|