kernel-elastic-autoencoder 3.4.0__tar.gz → 3.5.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.4.0
3
+ Version: 3.5.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.4.0"
3
+ version = "3.5.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" }
@@ -52,7 +52,6 @@ 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,
56
55
  **kwargs,
57
56
  ) -> Iterable[str]:
58
57
  """Completes each conditioned input sequence.
@@ -131,6 +130,7 @@ class Pipeline:
131
130
  beam_size: int,
132
131
  sequences: list[str],
133
132
  conditions: list[list[float]] | torch.Tensor,
133
+ hp_alpha: float = 0.5,
134
134
  **kwargs,
135
135
  ) -> Iterable[str]:
136
136
  """Completes each conditioned input sequence, using the beam search strategy.
@@ -144,6 +144,7 @@ class Pipeline:
144
144
  beam_size: Beam size of first step.
145
145
  sequences: List of text sequences to complete.
146
146
  conditions: List of condition value lists per batch.
147
+ hp_alpha: Strength of length normalization.
147
148
  **kwargs: Additional keyword arguments passed to Tokenizer.encode.
148
149
 
149
150
  Returns:
@@ -265,9 +266,7 @@ class Pipeline:
265
266
  top_prob_inds = (
266
267
  (
267
268
  top_probs.sum(dim=-1)
268
- / torch.sqrt(grouped_ids != self.tokenizer.pad_token_id)
269
- .to(torch.long)
270
- .sum(dim=-1)
269
+ / torch.pow((grouped_ids != self.tokenizer.pad_token_id).to(torch.long).sum(dim=-1), hp_alpha)
271
270
  )
272
271
  .topk(k=beam_size, dim=1)
273
272
  .indices.squeeze(-1)
@@ -297,9 +296,7 @@ class Pipeline:
297
296
  top_prob_inds = (
298
297
  (
299
298
  top_probs.sum(dim=-1)
300
- / torch.sqrt(grouped_ids != self.tokenizer.pad_token_id)
301
- .to(torch.long)
302
- .sum(dim=-1)
299
+ / torch.pow((grouped_ids != self.tokenizer.pad_token_id).to(torch.long).sum(dim=-1), hp_alpha)
303
300
  )
304
301
  .topk(k=1, dim=1)
305
302
  .indices.squeeze(-1)
@@ -194,7 +194,7 @@ class Trainer:
194
194
 
195
195
  curr_epoch += 1
196
196
  accelerator.wait_for_everyone()
197
- accelerator.save_state(checkpoint)
197
+ accelerator.save_state(checkpoint, safe_serialization=False, save_on_each_node=True)
198
198
  if accelerator.is_main_process:
199
199
  os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
200
200
  all_contexts.append(cb_ctx)