kernel-elastic-autoencoder 3.3.5__tar.gz → 3.3.6__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.3.6}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/pipeline.py +58 -21
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/training.py +4 -2
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/README.md +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/losses.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.3.5 → kernel_elastic_autoencoder-3.3.6}/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.
|
|
3
|
+
version = "3.3.6"
|
|
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" }
|
|
@@ -108,7 +108,7 @@ class Pipeline:
|
|
|
108
108
|
.indices.squeeze(-1)
|
|
109
109
|
.to(torch.long)
|
|
110
110
|
)
|
|
111
|
-
batches_completed |= new_toks.
|
|
111
|
+
batches_completed |= new_toks.flatten() == self.tokenizer.eos_token_id
|
|
112
112
|
new_toks = torch.where(
|
|
113
113
|
batches_completed.unsqueeze(-1),
|
|
114
114
|
self.tokenizer.pad_token_id,
|
|
@@ -118,9 +118,7 @@ class Pipeline:
|
|
|
118
118
|
token_mask = torch.cat(
|
|
119
119
|
[
|
|
120
120
|
token_mask,
|
|
121
|
-
|
|
122
|
-
input_ids.size(0), 1, dtype=torch.bool, device=self.device
|
|
123
|
-
),
|
|
121
|
+
batches_completed.unsqueeze(-1),
|
|
124
122
|
],
|
|
125
123
|
dim=1,
|
|
126
124
|
)
|
|
@@ -136,10 +134,9 @@ class Pipeline:
|
|
|
136
134
|
) -> Iterable[str]:
|
|
137
135
|
"""Completes each conditioned input sequence, using the beam search strategy.
|
|
138
136
|
|
|
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.
|
|
137
|
+
The model completes each sequence in the provided list using decoder-only inference. The beam search
|
|
138
|
+
decoding algorithm is used. On the last step, instead of returning beam_size candidates per batch, the
|
|
139
|
+
candidate with the highest length-normalized sum of log odds is selected for each batch.
|
|
143
140
|
|
|
144
141
|
Args:
|
|
145
142
|
latents: Tensor of dimension (B, P * E) containing latent vectors for the batch.
|
|
@@ -212,11 +209,11 @@ class Pipeline:
|
|
|
212
209
|
conds_embed = conds_embed.repeat_interleave(beam_size, dim=0)
|
|
213
210
|
latents = latents.repeat_interleave(beam_size, dim=0)
|
|
214
211
|
token_mask = token_mask.repeat_interleave(beam_size, dim=0)
|
|
215
|
-
batches_completed |= new_toks.
|
|
212
|
+
batches_completed |= new_toks.flatten() == self.tokenizer.eos_token_id
|
|
216
213
|
token_mask = torch.cat(
|
|
217
214
|
[
|
|
218
215
|
token_mask,
|
|
219
|
-
|
|
216
|
+
batches_completed.unsqueeze(-1),
|
|
220
217
|
],
|
|
221
218
|
dim=1,
|
|
222
219
|
)
|
|
@@ -230,41 +227,81 @@ class Pipeline:
|
|
|
230
227
|
condition_embeddings=conds_embed,
|
|
231
228
|
token_mask=token_mask,
|
|
232
229
|
)
|
|
230
|
+
odds = logits.log_softmax(dim=-1)
|
|
233
231
|
new_toks = (
|
|
234
|
-
torch.topk(
|
|
232
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
235
233
|
.indices.flatten()
|
|
236
234
|
.unsqueeze(-1)
|
|
237
235
|
.to(torch.long)
|
|
238
236
|
)
|
|
239
237
|
new_probs = (
|
|
240
|
-
torch.topk(
|
|
238
|
+
torch.topk(odds[:, -1:], k=beam_size, dim=-1)
|
|
239
|
+
.values.flatten()
|
|
240
|
+
.unsqueeze(-1)
|
|
241
241
|
)
|
|
242
242
|
new_toks = torch.where(
|
|
243
|
-
batches_completed.unsqueeze(-1),
|
|
243
|
+
batches_completed.repeat_interleave(beam_size, dim=0).unsqueeze(-1),
|
|
244
244
|
self.tokenizer.pad_token_id,
|
|
245
245
|
new_toks,
|
|
246
246
|
)
|
|
247
247
|
new_probs = torch.where(
|
|
248
|
-
batches_completed.unsqueeze(-1),
|
|
248
|
+
batches_completed.repeat_interleave(beam_size, dim=0).unsqueeze(-1),
|
|
249
249
|
0.0,
|
|
250
250
|
new_probs,
|
|
251
251
|
)
|
|
252
|
-
|
|
253
|
-
|
|
252
|
+
|
|
253
|
+
candidate_ids = torch.cat(
|
|
254
|
+
[input_ids.repeat_interleave(beam_size, dim=0), new_toks], dim=1
|
|
255
|
+
)
|
|
256
|
+
candidate_probs = torch.cat(
|
|
257
|
+
[input_probs.repeat_interleave(beam_size, dim=0), new_probs], dim=1
|
|
258
|
+
)
|
|
259
|
+
top_probs = candidate_probs.view(
|
|
260
|
+
candidate_probs.size(0) // (beam_size**2), beam_size**2, -1
|
|
261
|
+
)
|
|
262
|
+
grouped_ids = candidate_ids.view(top_probs.size(0), top_probs.size(1), -1)
|
|
263
|
+
top_prob_inds = (
|
|
264
|
+
(
|
|
265
|
+
top_probs.sum(dim=-1)
|
|
266
|
+
* (grouped_ids != self.tokenizer.eos_token_id)
|
|
267
|
+
.to(torch.long)
|
|
268
|
+
.sum(dim=-1)
|
|
269
|
+
)
|
|
270
|
+
.topk(k=beam_size, dim=1)
|
|
271
|
+
.indices.squeeze(-1)
|
|
272
|
+
)
|
|
273
|
+
input_ids = grouped_ids[
|
|
274
|
+
torch.arange(top_probs.size(0)).unsqueeze(-1).repeat(1, beam_size),
|
|
275
|
+
top_prob_inds,
|
|
276
|
+
].view(input_ids.size(0), -1)
|
|
277
|
+
input_probs = top_probs[
|
|
278
|
+
torch.arange(top_probs.size(0)).unsqueeze(-1).repeat(1, beam_size),
|
|
279
|
+
top_prob_inds,
|
|
280
|
+
].view(input_ids.size(0), -1)
|
|
281
|
+
|
|
282
|
+
batches_completed |= (
|
|
283
|
+
input_ids[:, -1:].flatten() == self.tokenizer.eos_token_id
|
|
284
|
+
)
|
|
254
285
|
token_mask = torch.cat(
|
|
255
286
|
[
|
|
256
287
|
token_mask,
|
|
257
|
-
|
|
258
|
-
input_ids.size(0), 1, dtype=torch.bool, device=self.device
|
|
259
|
-
),
|
|
288
|
+
batches_completed.unsqueeze(-1),
|
|
260
289
|
],
|
|
261
290
|
dim=1,
|
|
262
291
|
)
|
|
263
|
-
batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
264
292
|
|
|
265
293
|
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
294
|
grouped_ids = input_ids.view(top_probs.shape[0], beam_size, -1)
|
|
295
|
+
top_prob_inds = (
|
|
296
|
+
(
|
|
297
|
+
top_probs.sum(dim=-1)
|
|
298
|
+
* (grouped_ids != self.tokenizer.eos_token_id)
|
|
299
|
+
.to(torch.long)
|
|
300
|
+
.sum(dim=-1)
|
|
301
|
+
)
|
|
302
|
+
.topk(k=1, dim=1)
|
|
303
|
+
.indices.squeeze(-1)
|
|
304
|
+
)
|
|
268
305
|
winning_ids = grouped_ids[torch.arange(top_probs.size(0)), top_prob_inds]
|
|
269
306
|
return self.tokenizer.decode(winning_ids, skip_special_tokens=True)
|
|
270
307
|
|
|
@@ -94,13 +94,15 @@ class Trainer:
|
|
|
94
94
|
dataset_train,
|
|
95
95
|
batch_size=self.config_typed.common.batch_size,
|
|
96
96
|
pin_memory=True,
|
|
97
|
-
num_workers=
|
|
97
|
+
num_workers=4,
|
|
98
|
+
shuffle=True,
|
|
98
99
|
)
|
|
99
100
|
dataloader_test = torch.utils.data.DataLoader(
|
|
100
101
|
dataset_test,
|
|
101
102
|
batch_size=self.config_typed.common.batch_size,
|
|
102
103
|
pin_memory=True,
|
|
103
|
-
num_workers=
|
|
104
|
+
num_workers=4,
|
|
105
|
+
shuffle=True,
|
|
104
106
|
)
|
|
105
107
|
curr_epoch = 0
|
|
106
108
|
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|