kernel-elastic-autoencoder 1.0.0__tar.gz → 2.0.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-1.0.0 → kernel_elastic_autoencoder-2.0.0}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/pyproject.toml +4 -2
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/losses.py +1 -1
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/model.py +0 -44
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/pipeline.py +5 -5
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/sample.py +17 -5
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/README.md +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/__init__.py +14 -14
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/collate.py +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-2.0.0}/src/kernel_elastic_autoencoder/training.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernel_elastic_autoencoder"
|
|
3
|
-
version = "
|
|
3
|
+
version = "2.0.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" }
|
|
@@ -18,6 +18,9 @@ dependencies = [
|
|
|
18
18
|
[tool.poetry]
|
|
19
19
|
packages = [{ include = "kernel_elastic_autoencoder", from = "src" }]
|
|
20
20
|
|
|
21
|
+
[tool.ruff.lint]
|
|
22
|
+
ignore = ["B008"]
|
|
23
|
+
|
|
21
24
|
[tool.semantic_release]
|
|
22
25
|
version_toml = ["pyproject.toml:project.version"]
|
|
23
26
|
commit_parser = "conventional"
|
|
@@ -40,7 +43,6 @@ dev = [
|
|
|
40
43
|
"ruff (>=0.15.20,<0.17.0)",
|
|
41
44
|
"pytest-cov (>=7.1.0,<8.0.0)",
|
|
42
45
|
"pytest-xdist (>=3.8.0,<4.0.0)",
|
|
43
|
-
"ty (>=0.0.62,<0.0.64)",
|
|
44
46
|
"python-semantic-release (>=10.6.1,<11.0.0)",
|
|
45
47
|
"pdoc (>=16.0.0,<17.0.0)",
|
|
46
48
|
]
|
|
@@ -121,7 +121,7 @@ class Loss(nn.Module):
|
|
|
121
121
|
/ (2 * (self.hp_sigma**2))
|
|
122
122
|
).sum()
|
|
123
123
|
loss = self.hp_lambda * (
|
|
124
|
-
torch.tensor(
|
|
124
|
+
torch.tensor(1, dtype=latents.dtype, device=latents.device)
|
|
125
125
|
- ((1 / (latents.size(0) * self.kernel_dist_size)) * kernel_pairwise_sum)
|
|
126
126
|
)
|
|
127
127
|
return loss
|
|
@@ -1,9 +1,6 @@
|
|
|
1
|
-
from pathlib import Path
|
|
2
|
-
from typing import Any
|
|
3
1
|
|
|
4
2
|
import torch
|
|
5
3
|
from huggingface_hub import PyTorchModelHubMixin
|
|
6
|
-
from huggingface_hub.hub_mixin import T, DataclassInstance
|
|
7
4
|
from torch import nn
|
|
8
5
|
|
|
9
6
|
from kernel_elastic_autoencoder.config import ModelConfig
|
|
@@ -78,47 +75,6 @@ class Model(
|
|
|
78
75
|
)
|
|
79
76
|
"""Model decoder module."""
|
|
80
77
|
|
|
81
|
-
@classmethod
|
|
82
|
-
def from_pretrained(
|
|
83
|
-
cls: type[T],
|
|
84
|
-
pretrained_model_name_or_path: str | Path,
|
|
85
|
-
*,
|
|
86
|
-
force_download: bool = False,
|
|
87
|
-
token: str | bool | None = None,
|
|
88
|
-
cache_dir: str | Path | None = None,
|
|
89
|
-
local_files_only: bool = False,
|
|
90
|
-
revision: str | None = None,
|
|
91
|
-
**model_kwargs,
|
|
92
|
-
) -> T:
|
|
93
|
-
return cls.from_pretrained(
|
|
94
|
-
pretrained_model_name_or_path,
|
|
95
|
-
force_download=force_download,
|
|
96
|
-
token=token,
|
|
97
|
-
cache_dir=cache_dir,
|
|
98
|
-
local_files_only=local_files_only,
|
|
99
|
-
revision=revision,
|
|
100
|
-
**model_kwargs,
|
|
101
|
-
)
|
|
102
|
-
|
|
103
|
-
def save_pretrained(
|
|
104
|
-
self,
|
|
105
|
-
save_directory: str | Path,
|
|
106
|
-
*,
|
|
107
|
-
config: dict | DataclassInstance | None = None,
|
|
108
|
-
repo_id: str | None = None,
|
|
109
|
-
push_to_hub: bool = False,
|
|
110
|
-
model_card_kwargs: dict[str, Any] | None = None,
|
|
111
|
-
**push_to_hub_kwargs,
|
|
112
|
-
) -> str | None:
|
|
113
|
-
return super().save_pretrained(
|
|
114
|
-
save_directory,
|
|
115
|
-
config=config,
|
|
116
|
-
repo_id=repo_id,
|
|
117
|
-
push_to_hub=push_to_hub,
|
|
118
|
-
model_card_kwargs=model_card_kwargs,
|
|
119
|
-
**push_to_hub_kwargs,
|
|
120
|
-
)
|
|
121
|
-
|
|
122
78
|
def forward(
|
|
123
79
|
self,
|
|
124
80
|
input_ids: torch.Tensor,
|
|
@@ -31,7 +31,7 @@ class Completion(BaseModel):
|
|
|
31
31
|
"""Tensor containing generated condition masks for reuse."""
|
|
32
32
|
|
|
33
33
|
|
|
34
|
-
class Pipeline[
|
|
34
|
+
class Pipeline[T]:
|
|
35
35
|
"""User-facing pipeline for inference.
|
|
36
36
|
|
|
37
37
|
Defines an easy-to-use API for decoder-only inference with a pretrained model.
|
|
@@ -51,7 +51,7 @@ class Pipeline[_T]:
|
|
|
51
51
|
self,
|
|
52
52
|
model: Model,
|
|
53
53
|
tokenizer: Tokenizer,
|
|
54
|
-
collator: Callable[[ModelConfig, Tokenizer], Collator[
|
|
54
|
+
collator: Callable[[ModelConfig, Tokenizer], Collator[T]] = DataframeCollator,
|
|
55
55
|
sampler: Callable[[Tokenizer], Sampler] = Top1Sampler,
|
|
56
56
|
device: torch.device | None = None,
|
|
57
57
|
) -> None:
|
|
@@ -87,7 +87,7 @@ class Pipeline[_T]:
|
|
|
87
87
|
def _completion_entry(
|
|
88
88
|
self,
|
|
89
89
|
latents: torch.Tensor,
|
|
90
|
-
dataset:
|
|
90
|
+
dataset: T,
|
|
91
91
|
seq_feature: str,
|
|
92
92
|
cond_features: list[str],
|
|
93
93
|
device: torch.device | None = None,
|
|
@@ -148,7 +148,7 @@ class Pipeline[_T]:
|
|
|
148
148
|
self, intermediate: CompletionIntermediate, **kwargs
|
|
149
149
|
) -> Completion:
|
|
150
150
|
return Completion(
|
|
151
|
-
outputs=self.
|
|
151
|
+
outputs=self.sampler(
|
|
152
152
|
intermediate.input_ids, skip_special_tokens=True
|
|
153
153
|
),
|
|
154
154
|
condition_embeddings=intermediate.condition_embeddings,
|
|
@@ -158,7 +158,7 @@ class Pipeline[_T]:
|
|
|
158
158
|
def completion(
|
|
159
159
|
self,
|
|
160
160
|
latents: torch.Tensor,
|
|
161
|
-
dataset:
|
|
161
|
+
dataset: T,
|
|
162
162
|
seq_feature: str,
|
|
163
163
|
cond_features: list[str],
|
|
164
164
|
device: torch.device | None = None,
|
|
@@ -12,25 +12,34 @@ class Sampler(ABC):
|
|
|
12
12
|
Provides a specification for sampling text sequences from logits.
|
|
13
13
|
"""
|
|
14
14
|
|
|
15
|
-
def __init__(self, tokenizer: Tokenizer):
|
|
15
|
+
def __init__(self, tokenizer: Tokenizer, **kwargs):
|
|
16
16
|
"""Instantiates a Sampler object.
|
|
17
17
|
|
|
18
18
|
Args:
|
|
19
19
|
tokenizer: Object implementing the Tokenizer protocol.
|
|
20
|
+
**kwargs: Sampler-specific keyword args.
|
|
20
21
|
"""
|
|
21
22
|
self.tokenizer = tokenizer
|
|
22
23
|
"""Object implementing the Tokenizer protocol."""
|
|
23
24
|
|
|
25
|
+
@abstractmethod
|
|
24
26
|
def __call__(
|
|
25
27
|
self, logits: torch.Tensor, skip_special_tokens: bool, **kwargs
|
|
26
28
|
) -> Iterable[str]:
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
29
|
+
"""Interface method for implementing the last sampling step where outputs are decoded through the tokenizer.
|
|
30
|
+
|
|
31
|
+
Args:
|
|
32
|
+
logits: Tensor of dimension (B, S) containing IDs at the last step of inference.
|
|
33
|
+
**kwargs: Keyword arguments.
|
|
34
|
+
|
|
35
|
+
Returns:
|
|
36
|
+
Iterable[str]: Iterable of strings decoded by the tokenizer.
|
|
37
|
+
"""
|
|
38
|
+
...
|
|
30
39
|
|
|
31
40
|
@abstractmethod
|
|
32
41
|
def sample_ids(self, logits: torch.Tensor, **kwargs) -> torch.Tensor:
|
|
33
|
-
"""Interface method for implementing index sampling from logits.
|
|
42
|
+
"""Interface method for implementing index sampling from logits in intermediate steps.
|
|
34
43
|
|
|
35
44
|
Args:
|
|
36
45
|
logits: Tensor of dimension (B, S, L) containing logits.
|
|
@@ -43,6 +52,9 @@ class Sampler(ABC):
|
|
|
43
52
|
|
|
44
53
|
|
|
45
54
|
class Top1Sampler(Sampler):
|
|
55
|
+
def __call__(self, logits: torch.Tensor, skip_special_tokens: bool, **kwargs) -> Iterable[str]:
|
|
56
|
+
return self.tokenizer.decode(logits, skip_special_tokens=skip_special_tokens)
|
|
57
|
+
|
|
46
58
|
def sample_ids(self, logits: torch.Tensor, **kwargs) -> torch.Tensor:
|
|
47
59
|
"""Implementation of index sampling from logits choosing the highest-probability token.
|
|
48
60
|
|
|
File without changes
|
|
@@ -19,26 +19,26 @@ from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
|
19
19
|
from kernel_elastic_autoencoder.training import Trainer, TrainerCallback
|
|
20
20
|
|
|
21
21
|
__all__ = [
|
|
22
|
+
"Collated",
|
|
23
|
+
"Collator",
|
|
24
|
+
"Completion",
|
|
25
|
+
"DataframeCollator",
|
|
22
26
|
"ExperimentConfig",
|
|
23
|
-
"
|
|
27
|
+
"Loss",
|
|
28
|
+
"Model",
|
|
24
29
|
"ModelCommonConfig",
|
|
25
|
-
"
|
|
26
|
-
"ModelEncoderConfig",
|
|
30
|
+
"ModelConfig",
|
|
27
31
|
"ModelDecoderConfig",
|
|
28
|
-
"
|
|
29
|
-
"
|
|
30
|
-
"TrainingHyperparameterConfig",
|
|
31
|
-
"TrainingOptimizerConfig",
|
|
32
|
+
"ModelEncoderConfig",
|
|
33
|
+
"ModelInputConfig",
|
|
32
34
|
"Pipeline",
|
|
33
|
-
"Completion",
|
|
34
|
-
"Model",
|
|
35
|
-
"Loss",
|
|
36
|
-
"Tokenizer",
|
|
37
|
-
"Collator",
|
|
38
|
-
"Collated",
|
|
39
|
-
"DataframeCollator",
|
|
40
35
|
"Sampler",
|
|
36
|
+
"Tokenizer",
|
|
41
37
|
"Top1Sampler",
|
|
42
38
|
"Trainer",
|
|
43
39
|
"TrainerCallback",
|
|
40
|
+
"TrainingCommonConfig",
|
|
41
|
+
"TrainingConfig",
|
|
42
|
+
"TrainingHyperparameterConfig",
|
|
43
|
+
"TrainingOptimizerConfig",
|
|
44
44
|
]
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|