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.
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/pyproject.toml +1 -1
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/pipeline.py +4 -7
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/training.py +1 -1
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/README.md +0 -0
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/losses.py +0 -0
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.4.0 → kernel_elastic_autoencoder-3.5.0}/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
|
+
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.
|
|
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.
|
|
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)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|