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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 3.3.5
3
+ Version: 3.4.0
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.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
- self.register_buffer("_loc", torch.zeros(self.pooling_dim * self.embedding_dim))
65
- self.register_buffer("_cov", torch.eye(self.pooling_dim * self.embedding_dim))
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
- mvn = torch.distributions.MultivariateNormal(
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.squeeze(-1) == self.tokenizer.eos_token_id
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
- torch.zeros(
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. 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.
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(logits[:, -1:], k=beam_size, dim=-1)
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(logits[:, -1:], k=beam_size, dim=-1)
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.squeeze(-1) == self.tokenizer.eos_token_id
214
+ batches_completed |= new_toks.flatten() == self.tokenizer.eos_token_id
216
215
  token_mask = torch.cat(
217
216
  [
218
217
  token_mask,
219
- torch.zeros(input_ids.size(0), 1, dtype=torch.bool, device=self.device),
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(logits[:, -1:], k=1, dim=-1)
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(logits[:, -1:], k=1, dim=-1).values.flatten().unsqueeze(-1)
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
- input_ids = torch.cat([input_ids, new_toks], dim=1)
253
- input_probs = torch.cat([input_probs, new_probs], dim=1)
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
- torch.zeros(
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=0,
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=0,
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
- test_loss = torch.cat([test_loss, loss.detach().unsqueeze(-1)], dim=0)
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
  )