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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 3.3.5
3
+ Version: 3.3.6
4
4
  Summary: Implementation of Kernel-Elastic Autoencoder for Molecular Design (https://doi.org/10.1093/pnasnexus/pgae168)
5
5
  License: MIT
6
6
  Author: Felix Rotter-McCartney
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "kernel_elastic_autoencoder"
3
- version = "3.3.5"
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.squeeze(-1) == self.tokenizer.eos_token_id
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
- torch.zeros(
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. On the
140
- first step, the top `beam_size` first tokens are chosen for each batch. Subsequent tokens are
141
- sampled greedily in parallel for every subsequence. Before returning output, the subsequence with
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.squeeze(-1) == self.tokenizer.eos_token_id
212
+ batches_completed |= new_toks.flatten() == self.tokenizer.eos_token_id
216
213
  token_mask = torch.cat(
217
214
  [
218
215
  token_mask,
219
- torch.zeros(input_ids.size(0), 1, dtype=torch.bool, device=self.device),
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(logits[:, -1:], k=1, dim=-1)
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(logits[:, -1:], k=1, dim=-1).values.flatten().unsqueeze(-1)
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
- input_ids = torch.cat([input_ids, new_toks], dim=1)
253
- input_probs = torch.cat([input_probs, new_probs], dim=1)
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
- torch.zeros(
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=0,
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=0,
104
+ num_workers=4,
105
+ shuffle=True,
104
106
  )
105
107
  curr_epoch = 0
106
108