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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 3.3.6
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.6"
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.
@@ -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(logits[:, -1:], k=beam_size, dim=-1)
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(logits[:, -1:], k=beam_size, dim=-1)
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
- * (grouped_ids != self.tokenizer.eos_token_id)
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
- * (grouped_ids != self.tokenizer.eos_token_id)
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
- 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)
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
  )