kernel-elastic-autoencoder 1.0.0__tar.gz → 3.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-3.0.0}/PKG-INFO +2 -2
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/pyproject.toml +5 -3
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/__init__.py +11 -20
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/losses.py +1 -1
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/model.py +9 -46
- kernel_elastic_autoencoder-3.0.0/src/kernel_elastic_autoencoder/pipeline.py +161 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -68
- kernel_elastic_autoencoder-3.0.0/src/kernel_elastic_autoencoder/training.py +153 -0
- kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/collate.py +0 -164
- kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/pipeline.py +0 -203
- kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/sample.py +0 -58
- kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/training.py +0 -386
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/README.md +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: kernel_elastic_autoencoder
|
|
3
|
-
Version:
|
|
3
|
+
Version: 3.0.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
|
|
@@ -11,8 +11,8 @@ Classifier: Programming Language :: Python :: 3
|
|
|
11
11
|
Classifier: Programming Language :: Python :: 3.12
|
|
12
12
|
Classifier: Programming Language :: Python :: 3.13
|
|
13
13
|
Classifier: Programming Language :: Python :: 3.14
|
|
14
|
+
Requires-Dist: accelerate (>=1.14.0,<2.0.0)
|
|
14
15
|
Requires-Dist: huggingface-hub (>=1.22.0,<2.0.0)
|
|
15
|
-
Requires-Dist: pandas (>=3.0.3,<4.0.0)
|
|
16
16
|
Requires-Dist: pydantic (>=2.13.4,<3.0.0)
|
|
17
17
|
Requires-Dist: safetensors (>=0.8.0,<0.9.0)
|
|
18
18
|
Description-Content-Type: text/markdown
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernel_elastic_autoencoder"
|
|
3
|
-
version = "
|
|
3
|
+
version = "3.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" }
|
|
@@ -11,13 +11,16 @@ requires-python = ">=3.12,<3.15"
|
|
|
11
11
|
dependencies = [
|
|
12
12
|
"pydantic (>=2.13.4,<3.0.0)",
|
|
13
13
|
"huggingface-hub (>=1.22.0,<2.0.0)",
|
|
14
|
-
"pandas (>=3.0.3,<4.0.0)",
|
|
15
14
|
"safetensors (>=0.8.0,<0.9.0)",
|
|
15
|
+
"accelerate (>=1.14.0,<2.0.0)",
|
|
16
16
|
]
|
|
17
17
|
|
|
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
|
]
|
|
@@ -1,4 +1,3 @@
|
|
|
1
|
-
from kernel_elastic_autoencoder.collate import Collated, Collator, DataframeCollator
|
|
2
1
|
from kernel_elastic_autoencoder.config import (
|
|
3
2
|
ExperimentConfig,
|
|
4
3
|
ModelCommonConfig,
|
|
@@ -13,32 +12,24 @@ from kernel_elastic_autoencoder.config import (
|
|
|
13
12
|
)
|
|
14
13
|
from kernel_elastic_autoencoder.losses import Loss
|
|
15
14
|
from kernel_elastic_autoencoder.model import Model
|
|
16
|
-
from kernel_elastic_autoencoder.pipeline import
|
|
17
|
-
from kernel_elastic_autoencoder.sample import Sampler, Top1Sampler
|
|
15
|
+
from kernel_elastic_autoencoder.pipeline import Pipeline
|
|
18
16
|
from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
19
|
-
from kernel_elastic_autoencoder.training import Trainer
|
|
17
|
+
from kernel_elastic_autoencoder.training import Trainer
|
|
20
18
|
|
|
21
19
|
__all__ = [
|
|
22
20
|
"ExperimentConfig",
|
|
23
|
-
"
|
|
21
|
+
"Loss",
|
|
22
|
+
"Model",
|
|
24
23
|
"ModelCommonConfig",
|
|
25
|
-
"
|
|
26
|
-
"ModelEncoderConfig",
|
|
24
|
+
"ModelConfig",
|
|
27
25
|
"ModelDecoderConfig",
|
|
28
|
-
"
|
|
29
|
-
"
|
|
30
|
-
"TrainingHyperparameterConfig",
|
|
31
|
-
"TrainingOptimizerConfig",
|
|
26
|
+
"ModelEncoderConfig",
|
|
27
|
+
"ModelInputConfig",
|
|
32
28
|
"Pipeline",
|
|
33
|
-
"Completion",
|
|
34
|
-
"Model",
|
|
35
|
-
"Loss",
|
|
36
29
|
"Tokenizer",
|
|
37
|
-
"Collator",
|
|
38
|
-
"Collated",
|
|
39
|
-
"DataframeCollator",
|
|
40
|
-
"Sampler",
|
|
41
|
-
"Top1Sampler",
|
|
42
30
|
"Trainer",
|
|
43
|
-
"
|
|
31
|
+
"TrainingCommonConfig",
|
|
32
|
+
"TrainingConfig",
|
|
33
|
+
"TrainingHyperparameterConfig",
|
|
34
|
+
"TrainingOptimizerConfig",
|
|
44
35
|
]
|
|
@@ -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,5 @@
|
|
|
1
|
-
from pathlib import Path
|
|
2
|
-
from typing import Any
|
|
3
|
-
|
|
4
1
|
import torch
|
|
5
2
|
from huggingface_hub import PyTorchModelHubMixin
|
|
6
|
-
from huggingface_hub.hub_mixin import T, DataclassInstance
|
|
7
3
|
from torch import nn
|
|
8
4
|
|
|
9
5
|
from kernel_elastic_autoencoder.config import ModelConfig
|
|
@@ -78,47 +74,6 @@ class Model(
|
|
|
78
74
|
)
|
|
79
75
|
"""Model decoder module."""
|
|
80
76
|
|
|
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
77
|
def forward(
|
|
123
78
|
self,
|
|
124
79
|
input_ids: torch.Tensor,
|
|
@@ -218,5 +173,13 @@ class Model(
|
|
|
218
173
|
).to(torch.bool)
|
|
219
174
|
return self.decoder(current_output, latents, condition_embeddings, padding_mask)
|
|
220
175
|
|
|
221
|
-
def embed_conditions(self, conditions):
|
|
176
|
+
def embed_conditions(self, conditions: torch.Tensor) -> torch.Tensor:
|
|
177
|
+
"""Basic interface for the separate embedding of condition vectors.
|
|
178
|
+
|
|
179
|
+
Args:
|
|
180
|
+
conditions: Tensor of dimension (B, C) containing condition values for each sequence.
|
|
181
|
+
|
|
182
|
+
Returns:
|
|
183
|
+
torch.Tensor: Tensor of dimension (B, C, E) containing condition embeddings for each sequence.
|
|
184
|
+
"""
|
|
222
185
|
return self.encoder.embedding.conditional_embedding(conditions)
|
|
@@ -0,0 +1,161 @@
|
|
|
1
|
+
from collections.abc import Iterable
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
|
|
5
|
+
from kernel_elastic_autoencoder.model import Model
|
|
6
|
+
from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class Pipeline:
|
|
10
|
+
"""User-facing pipeline for inference.
|
|
11
|
+
|
|
12
|
+
Defines an easy-to-use API for decoder-only inference with a pretrained model.
|
|
13
|
+
|
|
14
|
+
Examples:
|
|
15
|
+
Wrapping a pretrained model and tokenizer:
|
|
16
|
+
>>> model = Model.from_pretrained("./checkpoint")
|
|
17
|
+
>>> tokenizer = MyTokenizer.from_pretrained("./checkpoint/tokenizer")
|
|
18
|
+
>>> pipe = Pipeline(model, tokenizer)
|
|
19
|
+
|
|
20
|
+
Getting a completion for sequences:
|
|
21
|
+
>>> compl = pipe.completion(latents, ["abc", "def", "ghi"], [[1.0, 0.5], [2.0, 1.0], [3.0, 1.5]])
|
|
22
|
+
>>> print(compl.outputs)
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
model: Model,
|
|
28
|
+
tokenizer: Tokenizer,
|
|
29
|
+
device: torch.device | None = None,
|
|
30
|
+
) -> None:
|
|
31
|
+
"""Instantiates a Pipeline object.
|
|
32
|
+
|
|
33
|
+
Args:
|
|
34
|
+
model: Pre-trained Model object for inference. Can be obtained from Model.from_pretrained.
|
|
35
|
+
tokenizer: Pre-configured Tokenizer object. The Tokenizer protocol supports tokenizers
|
|
36
|
+
inheriting from transformers.PreTrainedTokenizerBase, so such tokenizers may be loaded
|
|
37
|
+
from HuggingFace Hub.
|
|
38
|
+
device: Torch device used for inference.
|
|
39
|
+
"""
|
|
40
|
+
self.model = model.to(device)
|
|
41
|
+
"""Pre-trained Model object for inference. Moved to Pipeline.device, and placed in eval() mode."""
|
|
42
|
+
self.model.eval()
|
|
43
|
+
self.tokenizer = tokenizer
|
|
44
|
+
"""Pre-configured Tokenizer object."""
|
|
45
|
+
self.device = device
|
|
46
|
+
"""Torch device used for inference."""
|
|
47
|
+
|
|
48
|
+
def _ingest(
|
|
49
|
+
self,
|
|
50
|
+
sequences: list[str],
|
|
51
|
+
conditions: list[list[float]] | torch.Tensor,
|
|
52
|
+
device: torch.device | None,
|
|
53
|
+
**kwargs,
|
|
54
|
+
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
55
|
+
input_ids = self.tokenizer.encode(
|
|
56
|
+
seq=sequences,
|
|
57
|
+
padding=False,
|
|
58
|
+
max_length=self.model.config_typed.input.max_len,
|
|
59
|
+
add_special_tokens=False,
|
|
60
|
+
**kwargs,
|
|
61
|
+
)
|
|
62
|
+
conditions = torch.as_tensor(conditions, dtype=torch.float, device=device)
|
|
63
|
+
condition_mask = (
|
|
64
|
+
conditions != self.model.config_typed.common.padding_value
|
|
65
|
+
).to(torch.bool)
|
|
66
|
+
return input_ids, conditions, condition_mask
|
|
67
|
+
|
|
68
|
+
@torch.inference_mode()
|
|
69
|
+
def _completion_entry(
|
|
70
|
+
self,
|
|
71
|
+
latents: torch.Tensor,
|
|
72
|
+
sequences: list[str],
|
|
73
|
+
conditions: list[list[float]] | torch.Tensor,
|
|
74
|
+
device: torch.device | None = None,
|
|
75
|
+
**kwargs,
|
|
76
|
+
) -> dict[str, torch.Tensor]:
|
|
77
|
+
input_ids, conds, cond_mask = self._ingest(
|
|
78
|
+
sequences, conditions, device, **kwargs
|
|
79
|
+
)
|
|
80
|
+
conds_embed = self.model.embed_conditions(conds)
|
|
81
|
+
return {
|
|
82
|
+
"latents": latents,
|
|
83
|
+
"input_ids": input_ids,
|
|
84
|
+
"batches_completed": torch.zeros(input_ids.size(0), dtype=torch.bool),
|
|
85
|
+
"condition_embeddings": conds_embed,
|
|
86
|
+
"condition_mask": cond_mask,
|
|
87
|
+
}
|
|
88
|
+
|
|
89
|
+
@torch.inference_mode()
|
|
90
|
+
def _completion_step(
|
|
91
|
+
self, intermediate: dict[str, torch.Tensor]
|
|
92
|
+
) -> dict[str, torch.Tensor]:
|
|
93
|
+
intermediate["input_ids"] = torch.cat(
|
|
94
|
+
[
|
|
95
|
+
torch.full(
|
|
96
|
+
(intermediate["input_ids"].size(0), 1), self.tokenizer.bos_token_id
|
|
97
|
+
),
|
|
98
|
+
intermediate["input_ids"],
|
|
99
|
+
],
|
|
100
|
+
dim=1,
|
|
101
|
+
)
|
|
102
|
+
logits = self.model.decode(
|
|
103
|
+
current_output=intermediate["input_ids"],
|
|
104
|
+
latents=intermediate["latents"],
|
|
105
|
+
condition_embeddings=intermediate["condition_embeddings"],
|
|
106
|
+
token_mask=None,
|
|
107
|
+
condition_mask=intermediate["condition_mask"],
|
|
108
|
+
)
|
|
109
|
+
new_toks = (
|
|
110
|
+
torch.topk(logits[:, -1:], k=1, dim=-1).indices.squeeze(-1).to(torch.long)
|
|
111
|
+
)
|
|
112
|
+
intermediate["batches_completed"] |= (
|
|
113
|
+
new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
114
|
+
)
|
|
115
|
+
new_toks = torch.where(
|
|
116
|
+
intermediate["batches_completed"].unsqueeze(-1),
|
|
117
|
+
self.tokenizer.pad_token_id,
|
|
118
|
+
new_toks,
|
|
119
|
+
)
|
|
120
|
+
intermediate["input_ids"] = torch.cat([intermediate["input_ids"], new_toks], dim=1)
|
|
121
|
+
return intermediate
|
|
122
|
+
|
|
123
|
+
@torch.inference_mode()
|
|
124
|
+
def _completion_exit(self, intermediate: dict[str, torch.Tensor]) -> Iterable[str]:
|
|
125
|
+
return self.tokenizer.decode(intermediate["input_ids"], skip_special_tokens=True)
|
|
126
|
+
|
|
127
|
+
def completion(
|
|
128
|
+
self,
|
|
129
|
+
latents: torch.Tensor,
|
|
130
|
+
sequences: list[str],
|
|
131
|
+
conditions: list[list[float]] | torch.Tensor,
|
|
132
|
+
device: torch.device | None = None,
|
|
133
|
+
**kwargs,
|
|
134
|
+
) -> Iterable[str]:
|
|
135
|
+
"""Completes each conditioned input sequence.
|
|
136
|
+
|
|
137
|
+
The model completes each sequence in the provided list using decoder-only inference. Tokens are
|
|
138
|
+
sampled greedily.
|
|
139
|
+
|
|
140
|
+
Args:
|
|
141
|
+
latents: Tensor of dimension (B, P * E) containing latent vectors for the batch.
|
|
142
|
+
sequences: List of text sequences to complete.
|
|
143
|
+
conditions: List of condition value lists per batch.
|
|
144
|
+
device: Torch device used for inference.
|
|
145
|
+
**kwargs: Additional keyword arguments passed to Tokenizer.encode.
|
|
146
|
+
|
|
147
|
+
Returns:
|
|
148
|
+
Iterable[str]: List of completed sequences, stripped of special tokens.
|
|
149
|
+
"""
|
|
150
|
+
intermediate = self._completion_entry(
|
|
151
|
+
latents=latents,
|
|
152
|
+
sequences=sequences,
|
|
153
|
+
conditions=conditions,
|
|
154
|
+
device=device,
|
|
155
|
+
**kwargs,
|
|
156
|
+
)
|
|
157
|
+
while (
|
|
158
|
+
intermediate["input_ids"].size(1) < self.model.config_typed.input.max_len
|
|
159
|
+
) and (not intermediate["batches_completed"].all()):
|
|
160
|
+
intermediate = self._completion_step(intermediate=intermediate)
|
|
161
|
+
return self._completion_exit(intermediate=intermediate)
|
|
@@ -78,71 +78,3 @@ class Tokenizer(Protocol):
|
|
|
78
78
|
Tokenizer: Pretrained tokenizer.
|
|
79
79
|
"""
|
|
80
80
|
...
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
class _DummySlowTokenizer(Tokenizer):
|
|
84
|
-
def __init__(self, train: Iterable[str]):
|
|
85
|
-
self.pad_token = "_"
|
|
86
|
-
self.bos_token = "?"
|
|
87
|
-
self.eos_token = "!"
|
|
88
|
-
|
|
89
|
-
vocab = set[str]()
|
|
90
|
-
for b in train:
|
|
91
|
-
vocab = vocab.union(set(b))
|
|
92
|
-
|
|
93
|
-
special_tokens = [self.pad_token, self.bos_token, self.eos_token]
|
|
94
|
-
special_tokens.extend(list(vocab))
|
|
95
|
-
self.vocab = special_tokens
|
|
96
|
-
self.vocab_size = len(self.vocab)
|
|
97
|
-
|
|
98
|
-
self.pad_token_id = 0
|
|
99
|
-
self.bos_token_id = 1
|
|
100
|
-
self.eos_token_id = 2
|
|
101
|
-
|
|
102
|
-
def encode(
|
|
103
|
-
self,
|
|
104
|
-
seq: Iterable[str],
|
|
105
|
-
padding: bool,
|
|
106
|
-
max_length: int,
|
|
107
|
-
add_special_tokens: bool,
|
|
108
|
-
**kwargs,
|
|
109
|
-
) -> torch.Tensor:
|
|
110
|
-
ids_col = list[torch.Tensor]()
|
|
111
|
-
for b in seq:
|
|
112
|
-
toks = list(b)
|
|
113
|
-
if add_special_tokens:
|
|
114
|
-
toks.insert(0, self.bos_token)
|
|
115
|
-
toks.append(self.eos_token)
|
|
116
|
-
if padding:
|
|
117
|
-
toks += [self.pad_token] * (max_length - len(toks))
|
|
118
|
-
ids = torch.tensor([self.vocab.index(t) for t in toks], dtype=torch.long)
|
|
119
|
-
ids_col.append(ids)
|
|
120
|
-
return torch.stack(ids_col, dim=0)
|
|
121
|
-
|
|
122
|
-
def decode(
|
|
123
|
-
self, ids: torch.Tensor, skip_special_tokens: bool, **kwargs
|
|
124
|
-
) -> Iterable[str]:
|
|
125
|
-
if skip_special_tokens:
|
|
126
|
-
return [
|
|
127
|
-
str(
|
|
128
|
-
[
|
|
129
|
-
(
|
|
130
|
-
self.vocab[idx]
|
|
131
|
-
if idx
|
|
132
|
-
not in [
|
|
133
|
-
self.pad_token_id,
|
|
134
|
-
self.bos_token_id,
|
|
135
|
-
self.eos_token_id,
|
|
136
|
-
]
|
|
137
|
-
else ""
|
|
138
|
-
)
|
|
139
|
-
for idx in b
|
|
140
|
-
]
|
|
141
|
-
)
|
|
142
|
-
for b in ids
|
|
143
|
-
]
|
|
144
|
-
else:
|
|
145
|
-
return [str([self.vocab[idx] for idx in b]) for b in ids]
|
|
146
|
-
|
|
147
|
-
@classmethod
|
|
148
|
-
def from_pretrained(cls, pretrained_model_name_or_path: str | Path, **kwargs): ...
|
|
@@ -0,0 +1,153 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from collections.abc import Iterable
|
|
3
|
+
|
|
4
|
+
import torch
|
|
5
|
+
from accelerate import Accelerator
|
|
6
|
+
from accelerate.utils import tqdm
|
|
7
|
+
|
|
8
|
+
from kernel_elastic_autoencoder.config import TrainingConfig
|
|
9
|
+
from kernel_elastic_autoencoder.losses import Loss
|
|
10
|
+
from kernel_elastic_autoencoder.model import Model
|
|
11
|
+
from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class Trainer:
|
|
15
|
+
def __init__(
|
|
16
|
+
self,
|
|
17
|
+
config: dict | TrainingConfig,
|
|
18
|
+
):
|
|
19
|
+
"""Instantiates a Trainer object.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
config: Dictionary or TrainingConfig schema defining model parameters. Will be validated with TrainingConfig
|
|
23
|
+
regardless of input type.
|
|
24
|
+
"""
|
|
25
|
+
self.config = config
|
|
26
|
+
"""Configuration object for Hugging Face Hub compatible serialization. Not recommended to use, as it is
|
|
27
|
+
internal-use. Use Trainer.config_typed instead."""
|
|
28
|
+
config_typed = TrainingConfig.model_validate(config)
|
|
29
|
+
self.config_typed: TrainingConfig = config_typed
|
|
30
|
+
"""Type-validated config in a Pydantic TrainingConfig schema, recommended for public API use."""
|
|
31
|
+
|
|
32
|
+
def train(
|
|
33
|
+
self,
|
|
34
|
+
model: Model,
|
|
35
|
+
tokenizer: Tokenizer,
|
|
36
|
+
sequences: Iterable[str],
|
|
37
|
+
conditions: Iterable[Iterable[float]] | torch.Tensor,
|
|
38
|
+
train_split: float = 0.9,
|
|
39
|
+
checkpoint: str = "./checkpoint",
|
|
40
|
+
):
|
|
41
|
+
"""Trains a model. Optionally, resumes from an existing checkpoint.
|
|
42
|
+
|
|
43
|
+
Args:
|
|
44
|
+
model: Freshly instantiated model.
|
|
45
|
+
tokenizer: Pre-configured Tokenizer object. The Tokenizer protocol supports tokenizers
|
|
46
|
+
inheriting from transformers.PreTrainedTokenizerBase, so such tokenizers may be loaded
|
|
47
|
+
from HuggingFace Hub.
|
|
48
|
+
sequences: Iterable of text sequences.
|
|
49
|
+
conditions: Iterable of iterables of condition values per sequence.
|
|
50
|
+
train_split: Fraction of dataset used for training. Must be between 0 and 1.
|
|
51
|
+
checkpoint: Path of local checkpoint to be saved and/or resumed.
|
|
52
|
+
"""
|
|
53
|
+
accelerator = Accelerator()
|
|
54
|
+
|
|
55
|
+
loss_fn = Loss(
|
|
56
|
+
hp_lambda=self.config_typed.hyperparameters.hp_lambda,
|
|
57
|
+
hp_delta=self.config_typed.hyperparameters.hp_delta,
|
|
58
|
+
hp_sigma=self.config_typed.hyperparameters.hp_sigma,
|
|
59
|
+
kernel_dist_size=self.config_typed.hyperparameters.kernel_dist_size,
|
|
60
|
+
padding_idx=model.config_typed.common.padding_idx,
|
|
61
|
+
embedding_dim=model.config_typed.common.embedding_dim,
|
|
62
|
+
pooling_dim=model.config_typed.common.pooling_dim,
|
|
63
|
+
)
|
|
64
|
+
optimizer = self.config_typed.optimizer.optimizer_fn(
|
|
65
|
+
model.parameters(), **self.config_typed.optimizer.optimizer_params
|
|
66
|
+
)
|
|
67
|
+
scheduler = self.config_typed.optimizer.scheduler_fn(
|
|
68
|
+
optimizer, **self.config_typed.optimizer.scheduler_params
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
input_ids = tokenizer.encode(
|
|
72
|
+
seq=sequences,
|
|
73
|
+
padding=True,
|
|
74
|
+
max_length=model.config_typed.input.max_len,
|
|
75
|
+
add_special_tokens=True,
|
|
76
|
+
)
|
|
77
|
+
conditions = torch.as_tensor(conditions, dtype=torch.float)
|
|
78
|
+
token_mask = (input_ids != model.config_typed.common.padding_idx).to(torch.bool)
|
|
79
|
+
condition_mask = (conditions != model.config_typed.common.padding_value).to(
|
|
80
|
+
torch.bool
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
dataset = torch.utils.data.TensorDataset(
|
|
84
|
+
input_ids, conditions, token_mask, condition_mask
|
|
85
|
+
)
|
|
86
|
+
dataset_train, dataset_test = torch.utils.data.random_split(
|
|
87
|
+
dataset, [train_split, 1 - train_split]
|
|
88
|
+
)
|
|
89
|
+
dataloader_train = torch.utils.data.DataLoader(
|
|
90
|
+
dataset_train,
|
|
91
|
+
batch_size=self.config_typed.common.batch_size,
|
|
92
|
+
)
|
|
93
|
+
dataloader_test = torch.utils.data.DataLoader(
|
|
94
|
+
dataset_test,
|
|
95
|
+
batch_size=self.config_typed.common.batch_size,
|
|
96
|
+
)
|
|
97
|
+
curr_epoch = 0
|
|
98
|
+
|
|
99
|
+
model, optimizer, dataloader_train, dataloader_test, scheduler, curr_epoch = (
|
|
100
|
+
accelerator.prepare(
|
|
101
|
+
model,
|
|
102
|
+
optimizer,
|
|
103
|
+
dataloader_train,
|
|
104
|
+
dataloader_test,
|
|
105
|
+
scheduler,
|
|
106
|
+
curr_epoch,
|
|
107
|
+
)
|
|
108
|
+
)
|
|
109
|
+
accelerator.register_for_checkpointing(scheduler, curr_epoch)
|
|
110
|
+
|
|
111
|
+
if os.path.exists(checkpoint):
|
|
112
|
+
accelerator.load_state(checkpoint)
|
|
113
|
+
|
|
114
|
+
for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
|
|
115
|
+
model.train()
|
|
116
|
+
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
117
|
+
dataloader_train, desc=f"Epoch {epoch}, Train Batch"
|
|
118
|
+
):
|
|
119
|
+
optimizer.zero_grad()
|
|
120
|
+
prediction, prediction_noise, latents_noise = model(
|
|
121
|
+
input_ids, conditions, token_mask, condition_mask
|
|
122
|
+
)
|
|
123
|
+
loss = loss_fn(
|
|
124
|
+
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
125
|
+
)
|
|
126
|
+
accelerator.backward(loss)
|
|
127
|
+
optimizer.step()
|
|
128
|
+
|
|
129
|
+
model.eval()
|
|
130
|
+
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
131
|
+
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
132
|
+
):
|
|
133
|
+
with torch.no_grad():
|
|
134
|
+
prediction, prediction_noise, latents_noise = model(
|
|
135
|
+
input_ids,
|
|
136
|
+
conditions,
|
|
137
|
+
token_mask,
|
|
138
|
+
condition_mask,
|
|
139
|
+
)
|
|
140
|
+
loss = loss_fn(
|
|
141
|
+
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
142
|
+
)
|
|
143
|
+
|
|
144
|
+
scheduler.step(epoch)
|
|
145
|
+
|
|
146
|
+
accelerator.wait_for_everyone()
|
|
147
|
+
curr_epoch += 1
|
|
148
|
+
os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
|
|
149
|
+
model.save_pretrained(os.path.join(checkpoint, "dist/"))
|
|
150
|
+
self.config_typed.to_json(
|
|
151
|
+
os.path.join(checkpoint, "dist/train_config.json")
|
|
152
|
+
)
|
|
153
|
+
accelerator.save_state(checkpoint)
|