kernel-elastic-autoencoder 1.0.0__py3-none-any.whl
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/__init__.py +44 -0
- kernel_elastic_autoencoder/collate.py +164 -0
- kernel_elastic_autoencoder/config.py +250 -0
- kernel_elastic_autoencoder/layers.py +229 -0
- kernel_elastic_autoencoder/losses.py +127 -0
- kernel_elastic_autoencoder/model.py +222 -0
- kernel_elastic_autoencoder/pipeline.py +203 -0
- kernel_elastic_autoencoder/sample.py +58 -0
- kernel_elastic_autoencoder/tokenizer.py +148 -0
- kernel_elastic_autoencoder/training.py +386 -0
- kernel_elastic_autoencoder-1.0.0.dist-info/METADATA +44 -0
- kernel_elastic_autoencoder-1.0.0.dist-info/RECORD +13 -0
- kernel_elastic_autoencoder-1.0.0.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
from collections.abc import Iterable
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Protocol, Self
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class Tokenizer(Protocol):
|
|
9
|
+
"""Protocol to be implemented for tokenizers. Supports transformers.PreTrainedTokenizerBase.
|
|
10
|
+
|
|
11
|
+
Provides a specification for tokenizing text for model input, and recovering text from token indices.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
vocab_size: int
|
|
15
|
+
"""Number of tokens in the tokenizer's vocabulary."""
|
|
16
|
+
bos_token: str
|
|
17
|
+
"""Token marking the beginning of a sequence."""
|
|
18
|
+
bos_token_id: int
|
|
19
|
+
"""Index of token marking the beginning of a sequence."""
|
|
20
|
+
eos_token: str
|
|
21
|
+
"""Token marking the end of a sequence."""
|
|
22
|
+
eos_token_id: int
|
|
23
|
+
"""Index of token marking the end of a sequence."""
|
|
24
|
+
pad_token: str
|
|
25
|
+
"""Token representing padding."""
|
|
26
|
+
pad_token_id: int
|
|
27
|
+
"""Index of token representing padding."""
|
|
28
|
+
|
|
29
|
+
def encode(
|
|
30
|
+
self,
|
|
31
|
+
seq: Iterable[str],
|
|
32
|
+
padding: bool,
|
|
33
|
+
max_length: int,
|
|
34
|
+
add_special_tokens: bool,
|
|
35
|
+
**kwargs,
|
|
36
|
+
) -> torch.Tensor:
|
|
37
|
+
"""Encodes an Iterable of sequences to a tensor of indices. Optionally, adds special tokens according to a
|
|
38
|
+
template and pads the outputs to a fixed length.
|
|
39
|
+
|
|
40
|
+
Args:
|
|
41
|
+
seq: Iterable of text sequences to encode.
|
|
42
|
+
padding: Whether to pad the sequences to a fixed length.
|
|
43
|
+
max_length: Maximum length to which sequences are padded if padding is True.
|
|
44
|
+
add_special_tokens: Whether to add special tokens according to a template.
|
|
45
|
+
**kwargs: Keyword arguments.
|
|
46
|
+
|
|
47
|
+
Returns:
|
|
48
|
+
torch.Tensor: Tensor of dimension (B, S) containing vocabulary indices.
|
|
49
|
+
"""
|
|
50
|
+
...
|
|
51
|
+
|
|
52
|
+
def decode(
|
|
53
|
+
self, ids: torch.Tensor, skip_special_tokens: bool, **kwargs
|
|
54
|
+
) -> Iterable[str]:
|
|
55
|
+
"""Decodes a tensor of indices to an Iterable of sequences. Optionally, skips special tokens.
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
ids: Tensor of dimension (B, S) containing vocabulary indices to decode.
|
|
59
|
+
skip_special_tokens: Whether to skip decoding special tokens when constructing outputs.
|
|
60
|
+
**kwargs: Keyword arguments.
|
|
61
|
+
|
|
62
|
+
Returns:
|
|
63
|
+
Iterable[str]: Iterable of decoded sequences.
|
|
64
|
+
"""
|
|
65
|
+
...
|
|
66
|
+
|
|
67
|
+
@classmethod
|
|
68
|
+
def from_pretrained(
|
|
69
|
+
cls, pretrained_model_name_or_path: str | Path, **kwargs
|
|
70
|
+
) -> Self:
|
|
71
|
+
"""Loads a tokenizer from a path or remote name. Supports Hugging Face from_pretrained.
|
|
72
|
+
|
|
73
|
+
Args:
|
|
74
|
+
pretrained_model_name_or_path: Local path or remote name.
|
|
75
|
+
**kwargs: Keyword arguments.
|
|
76
|
+
|
|
77
|
+
Returns:
|
|
78
|
+
Tokenizer: Pretrained tokenizer.
|
|
79
|
+
"""
|
|
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,386 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from typing import Protocol
|
|
3
|
+
|
|
4
|
+
import torch
|
|
5
|
+
import torch.distributed as dist
|
|
6
|
+
from pydantic import BaseModel, ConfigDict
|
|
7
|
+
from torch import nn
|
|
8
|
+
from torch.nn.parallel import DistributedDataParallel as DDP
|
|
9
|
+
from torch.utils.data import DistributedSampler
|
|
10
|
+
|
|
11
|
+
from kernel_elastic_autoencoder.collate import Collated
|
|
12
|
+
from kernel_elastic_autoencoder.config import TrainingConfig
|
|
13
|
+
from kernel_elastic_autoencoder.losses import Loss
|
|
14
|
+
from kernel_elastic_autoencoder.model import Model
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _is_local():
|
|
18
|
+
return (
|
|
19
|
+
dist.is_torchelastic_launched() and (dist.get_rank() == 0)
|
|
20
|
+
) or not dist.is_torchelastic_launched()
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class TrainerCallbackCtx(BaseModel):
|
|
24
|
+
model_config = ConfigDict(arbitrary_types_allowed=True)
|
|
25
|
+
dist: bool
|
|
26
|
+
device_type: str
|
|
27
|
+
local_rank: int | None
|
|
28
|
+
model: Model | DDP
|
|
29
|
+
optimizer: torch.optim.Optimizer
|
|
30
|
+
scheduler: torch.optim.lr_scheduler.LRScheduler
|
|
31
|
+
dataloader_train: torch.utils.data.DataLoader
|
|
32
|
+
dataloader_test: torch.utils.data.DataLoader
|
|
33
|
+
|
|
34
|
+
epoch: int | None = None
|
|
35
|
+
rel_epoch: int | None = None
|
|
36
|
+
batch: int | None = None
|
|
37
|
+
train_loss: torch.Tensor | None = None
|
|
38
|
+
test_loss: torch.Tensor | None = None
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class TrainerCallback(Protocol):
|
|
42
|
+
"""Protocol to be implemented for callback classes to a Trainer.
|
|
43
|
+
|
|
44
|
+
On each hook, except before the Trainer is initialized, a TrainerCallbackCtx schema is passed with the appropriate
|
|
45
|
+
information passed.
|
|
46
|
+
"""
|
|
47
|
+
|
|
48
|
+
def before_init(self):
|
|
49
|
+
"""Called before training setup."""
|
|
50
|
+
...
|
|
51
|
+
|
|
52
|
+
def after_init(self, ctx: TrainerCallbackCtx):
|
|
53
|
+
"""Called after training setup."""
|
|
54
|
+
...
|
|
55
|
+
|
|
56
|
+
def before_epoch(self, ctx: TrainerCallbackCtx):
|
|
57
|
+
"""Called before each epoch."""
|
|
58
|
+
...
|
|
59
|
+
|
|
60
|
+
def before_train_batch(self, ctx: TrainerCallbackCtx):
|
|
61
|
+
"""Called before each forward pass of a single training batch."""
|
|
62
|
+
...
|
|
63
|
+
|
|
64
|
+
def after_train_batch(self, ctx: TrainerCallbackCtx):
|
|
65
|
+
"""Called after each forward pass of a single training batch."""
|
|
66
|
+
...
|
|
67
|
+
|
|
68
|
+
def before_test_batch(self, ctx: TrainerCallbackCtx):
|
|
69
|
+
"""Called before each forward pass of a single test batch."""
|
|
70
|
+
...
|
|
71
|
+
|
|
72
|
+
def after_test_batch(self, ctx: TrainerCallbackCtx):
|
|
73
|
+
"""Called after each forward pass of a single test batch."""
|
|
74
|
+
...
|
|
75
|
+
|
|
76
|
+
def after_epoch(self, ctx: TrainerCallbackCtx):
|
|
77
|
+
"""Called after each epoch."""
|
|
78
|
+
...
|
|
79
|
+
|
|
80
|
+
def after_training(self, ctx: TrainerCallbackCtx):
|
|
81
|
+
"""Called after training ends."""
|
|
82
|
+
...
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class TrainerDefaultCallback:
|
|
86
|
+
# TODO: Implement sensible default logging.
|
|
87
|
+
def before_init(self):
|
|
88
|
+
pass
|
|
89
|
+
|
|
90
|
+
def after_init(self, ctx: TrainerCallbackCtx):
|
|
91
|
+
pass
|
|
92
|
+
|
|
93
|
+
def before_epoch(self, ctx: TrainerCallbackCtx):
|
|
94
|
+
pass
|
|
95
|
+
|
|
96
|
+
def before_train_batch(self, ctx: TrainerCallbackCtx):
|
|
97
|
+
pass
|
|
98
|
+
|
|
99
|
+
def after_train_batch(self, ctx: TrainerCallbackCtx):
|
|
100
|
+
pass
|
|
101
|
+
|
|
102
|
+
def before_test_batch(self, ctx: TrainerCallbackCtx):
|
|
103
|
+
pass
|
|
104
|
+
|
|
105
|
+
def after_test_batch(self, ctx: TrainerCallbackCtx):
|
|
106
|
+
pass
|
|
107
|
+
|
|
108
|
+
def after_epoch(self, ctx: TrainerCallbackCtx):
|
|
109
|
+
pass
|
|
110
|
+
|
|
111
|
+
def after_training(self, ctx: TrainerCallbackCtx):
|
|
112
|
+
pass
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
class Trainer:
|
|
116
|
+
def __init__(
|
|
117
|
+
self,
|
|
118
|
+
config: dict | TrainingConfig,
|
|
119
|
+
callbacks: tuple[TrainerCallback] = (TrainerDefaultCallback(),),
|
|
120
|
+
):
|
|
121
|
+
"""Instantiates a Trainer object.
|
|
122
|
+
|
|
123
|
+
Args:
|
|
124
|
+
config: Dictionary or TrainingConfig schema defining model parameters. Will be validated with TrainingConfig
|
|
125
|
+
regardless of input type.
|
|
126
|
+
callbacks: Tuple of callback classes implementing TrainerCallback.
|
|
127
|
+
"""
|
|
128
|
+
self.config = config
|
|
129
|
+
"""Configuration object for Hugging Face Hub compatible serialization. Not recommended to use, as it is
|
|
130
|
+
internal-use. Use Trainer.config_typed instead."""
|
|
131
|
+
config_typed = TrainingConfig.model_validate(config)
|
|
132
|
+
self.config_typed: TrainingConfig = config_typed
|
|
133
|
+
"""Type-validated config in a Pydantic TrainingConfig schema, recommended for public API use."""
|
|
134
|
+
self.callbacks = callbacks
|
|
135
|
+
"""Tuple of callback classes implementing TrainerCallback."""
|
|
136
|
+
|
|
137
|
+
self._ctx: TrainerCallbackCtx
|
|
138
|
+
|
|
139
|
+
self._model: DDP | Model
|
|
140
|
+
self._dist: bool
|
|
141
|
+
self._loss_fn: nn.Module
|
|
142
|
+
self._optimizer: torch.optim.Optimizer
|
|
143
|
+
self._scheduler: torch.optim.lr_scheduler.LRScheduler
|
|
144
|
+
self._dataloader_train: torch.utils.data.DataLoader
|
|
145
|
+
self._dataloader_test: torch.utils.data.DataLoader
|
|
146
|
+
self._device_type: str
|
|
147
|
+
self._local_rank: int | None
|
|
148
|
+
|
|
149
|
+
def train(
|
|
150
|
+
self,
|
|
151
|
+
model: Model,
|
|
152
|
+
ds: Collated,
|
|
153
|
+
train_split: float = 0.9,
|
|
154
|
+
checkpoint: str = "./checkpoint",
|
|
155
|
+
):
|
|
156
|
+
"""Trains a model with a Collated dataset. Optionally resumes from an existing checkpoint.
|
|
157
|
+
|
|
158
|
+
Args:
|
|
159
|
+
model: Freshly instantiated model.
|
|
160
|
+
ds: Tensor dataset following the Collated schema.
|
|
161
|
+
train_split: Fraction of dataset used for training. Must be between 0 and 1.
|
|
162
|
+
checkpoint: Path of local checkpoint to be saved and/or resumed.
|
|
163
|
+
"""
|
|
164
|
+
next_epoch = self._setup(model, ds, train_split, checkpoint)
|
|
165
|
+
for epoch in range(next_epoch, self.config_typed.common.max_epochs):
|
|
166
|
+
self._ctx.rel_epoch = epoch - next_epoch
|
|
167
|
+
self._epoch(epoch, checkpoint)
|
|
168
|
+
self._ctx.epoch = None
|
|
169
|
+
self._ctx.rel_epoch = None
|
|
170
|
+
if _is_local():
|
|
171
|
+
for cb in self.callbacks:
|
|
172
|
+
cb.after_training(self._ctx)
|
|
173
|
+
|
|
174
|
+
def _setup(
|
|
175
|
+
self,
|
|
176
|
+
model: Model,
|
|
177
|
+
ds: Collated,
|
|
178
|
+
train_split: float,
|
|
179
|
+
checkpoint: str,
|
|
180
|
+
):
|
|
181
|
+
if _is_local():
|
|
182
|
+
for cb in self.callbacks:
|
|
183
|
+
cb.before_init()
|
|
184
|
+
self._model = model
|
|
185
|
+
self._dist = dist.is_torchelastic_launched()
|
|
186
|
+
if self._dist:
|
|
187
|
+
self._model, self._device_type, self._local_rank = self._setup_ddp()
|
|
188
|
+
else:
|
|
189
|
+
self._device_type, self._local_rank = self._setup_no_ddp()
|
|
190
|
+
self._loss_fn = self._setup_loss(model)
|
|
191
|
+
self._optimizer, self._scheduler = self._setup_optimizer()
|
|
192
|
+
self._dataloader_train, self._dataloader_test = self._setup_dataloaders(
|
|
193
|
+
ds, train_split
|
|
194
|
+
)
|
|
195
|
+
next_epoch = 0
|
|
196
|
+
if os.path.exists(checkpoint):
|
|
197
|
+
next_epoch = self._resume(checkpoint)
|
|
198
|
+
self._ctx = TrainerCallbackCtx(
|
|
199
|
+
dist=self._dist,
|
|
200
|
+
device_type=self._device_type,
|
|
201
|
+
local_rank=self._local_rank,
|
|
202
|
+
model=self._model,
|
|
203
|
+
optimizer=self._optimizer,
|
|
204
|
+
scheduler=self._scheduler,
|
|
205
|
+
dataloader_train=self._dataloader_train,
|
|
206
|
+
dataloader_test=self._dataloader_test,
|
|
207
|
+
)
|
|
208
|
+
if _is_local():
|
|
209
|
+
for cb in self.callbacks:
|
|
210
|
+
cb.after_init(self._ctx)
|
|
211
|
+
return next_epoch
|
|
212
|
+
|
|
213
|
+
def _setup_loss(self, model: Model):
|
|
214
|
+
_loss_fn = Loss(
|
|
215
|
+
hp_lambda=self.config_typed.hyperparameters.hp_lambda,
|
|
216
|
+
hp_delta=self.config_typed.hyperparameters.hp_delta,
|
|
217
|
+
hp_sigma=self.config_typed.hyperparameters.hp_sigma,
|
|
218
|
+
kernel_dist_size=self.config_typed.hyperparameters.kernel_dist_size,
|
|
219
|
+
padding_idx=model.config_typed.common.padding_idx,
|
|
220
|
+
embedding_dim=model.config_typed.common.embedding_dim,
|
|
221
|
+
pooling_dim=model.config_typed.common.pooling_dim,
|
|
222
|
+
)
|
|
223
|
+
return _loss_fn
|
|
224
|
+
|
|
225
|
+
def _setup_optimizer(self):
|
|
226
|
+
optimizer = self.config_typed.optimizer.optimizer_fn(
|
|
227
|
+
self._model.parameters(), **self.config_typed.optimizer.optimizer_params
|
|
228
|
+
)
|
|
229
|
+
scheduler = self.config_typed.optimizer.scheduler_fn(
|
|
230
|
+
optimizer, **self.config_typed.optimizer.scheduler_params
|
|
231
|
+
)
|
|
232
|
+
return optimizer, scheduler
|
|
233
|
+
|
|
234
|
+
def _setup_ddp(self) -> tuple[DDP, str, int]:
|
|
235
|
+
device_type, vendor_backend = self._get_backend()
|
|
236
|
+
dist.init_process_group(backend=vendor_backend)
|
|
237
|
+
local_rank = int(os.environ["LOCAL_RANK"])
|
|
238
|
+
model = DDP(self._model.to(local_rank))
|
|
239
|
+
return model, device_type, local_rank
|
|
240
|
+
|
|
241
|
+
def _setup_no_ddp(self):
|
|
242
|
+
device_type, _ = self._get_backend()
|
|
243
|
+
return device_type, None
|
|
244
|
+
|
|
245
|
+
def _setup_dataloaders(self, ds: Collated, train_split: float):
|
|
246
|
+
dataset = torch.utils.data.TensorDataset(
|
|
247
|
+
ds.input_ids, ds.conditions, ds.token_mask, ds.condition_mask
|
|
248
|
+
)
|
|
249
|
+
dataset_train, dataset_test = torch.utils.data.random_split(
|
|
250
|
+
dataset, [train_split, 1 - train_split]
|
|
251
|
+
)
|
|
252
|
+
if self._dist:
|
|
253
|
+
sampler_train = DistributedSampler(dataset_train)
|
|
254
|
+
sampler_test = DistributedSampler(dataset_test)
|
|
255
|
+
else:
|
|
256
|
+
sampler_train = torch.utils.data.RandomSampler(dataset_train)
|
|
257
|
+
sampler_test = torch.utils.data.SequentialSampler(dataset_test)
|
|
258
|
+
_dataloader_train = torch.utils.data.DataLoader(
|
|
259
|
+
dataset_train,
|
|
260
|
+
batch_size=self.config_typed.common.batch_size,
|
|
261
|
+
sampler=sampler_train,
|
|
262
|
+
pin_memory=True,
|
|
263
|
+
)
|
|
264
|
+
_dataloader_test = torch.utils.data.DataLoader(
|
|
265
|
+
dataset_test,
|
|
266
|
+
batch_size=self.config_typed.common.batch_size,
|
|
267
|
+
sampler=sampler_test,
|
|
268
|
+
pin_memory=True,
|
|
269
|
+
)
|
|
270
|
+
return _dataloader_train, _dataloader_test
|
|
271
|
+
|
|
272
|
+
def _resume(
|
|
273
|
+
self,
|
|
274
|
+
checkpoint: str,
|
|
275
|
+
):
|
|
276
|
+
train_state = torch.load(os.path.join(checkpoint, "train_state.pt"))
|
|
277
|
+
self._optimizer.load_state_dict(train_state["optimizer"])
|
|
278
|
+
self._scheduler.load_state_dict(train_state["scheduler"])
|
|
279
|
+
self._model.from_pretrained(checkpoint) # type: ignore
|
|
280
|
+
return train_state["next_epoch"]
|
|
281
|
+
|
|
282
|
+
def _epoch(
|
|
283
|
+
self,
|
|
284
|
+
epoch: int,
|
|
285
|
+
checkpoint: str,
|
|
286
|
+
):
|
|
287
|
+
self._ctx.epoch = epoch
|
|
288
|
+
if _is_local():
|
|
289
|
+
for cb in self.callbacks:
|
|
290
|
+
cb.before_epoch(self._ctx)
|
|
291
|
+
|
|
292
|
+
self._model.train()
|
|
293
|
+
self._ctx.batch = 0
|
|
294
|
+
for input_ids, conditions, token_mask, condition_mask in self._dataloader_train:
|
|
295
|
+
if _is_local():
|
|
296
|
+
for cb in self.callbacks:
|
|
297
|
+
cb.before_train_batch(self._ctx)
|
|
298
|
+
self._ctx.train_loss = self._batch_train(
|
|
299
|
+
input_ids, conditions, token_mask, condition_mask
|
|
300
|
+
)
|
|
301
|
+
if _is_local():
|
|
302
|
+
for cb in self.callbacks:
|
|
303
|
+
cb.after_train_batch(self._ctx)
|
|
304
|
+
self._ctx.batch += 1
|
|
305
|
+
|
|
306
|
+
self._model.eval()
|
|
307
|
+
self._ctx.batch = 0
|
|
308
|
+
for input_ids, conditions, token_mask, condition_mask in self._dataloader_test:
|
|
309
|
+
if _is_local():
|
|
310
|
+
for cb in self.callbacks:
|
|
311
|
+
cb.before_test_batch(self._ctx)
|
|
312
|
+
self._ctx.test_loss = self._batch_test(
|
|
313
|
+
input_ids, conditions, token_mask, condition_mask
|
|
314
|
+
)
|
|
315
|
+
if _is_local():
|
|
316
|
+
for cb in self.callbacks:
|
|
317
|
+
cb.after_test_batch(self._ctx)
|
|
318
|
+
self._ctx.batch += 1
|
|
319
|
+
|
|
320
|
+
self._end_epoch(epoch, checkpoint)
|
|
321
|
+
self._ctx.train_loss = None
|
|
322
|
+
self._ctx.test_loss = None
|
|
323
|
+
self._ctx.batch = None
|
|
324
|
+
if _is_local():
|
|
325
|
+
for cb in self.callbacks:
|
|
326
|
+
cb.after_epoch(self._ctx)
|
|
327
|
+
|
|
328
|
+
def _end_epoch(self, epoch: int, checkpoint: str):
|
|
329
|
+
self._scheduler.step(epoch)
|
|
330
|
+
self._save_state(checkpoint, epoch)
|
|
331
|
+
|
|
332
|
+
def _batch_train(self, input_ids, conditions, token_mask, condition_mask):
|
|
333
|
+
self._optimizer.zero_grad()
|
|
334
|
+
with torch.amp.autocast(self._device_type):
|
|
335
|
+
prediction, prediction_noise, latents_noise = self._model.forward(
|
|
336
|
+
input_ids, conditions, token_mask, condition_mask
|
|
337
|
+
)
|
|
338
|
+
loss = self._loss_fn(
|
|
339
|
+
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
340
|
+
)
|
|
341
|
+
loss.backward()
|
|
342
|
+
self._optimizer.step()
|
|
343
|
+
return loss.detach()
|
|
344
|
+
|
|
345
|
+
def _batch_test(self, input_ids, conditions, token_mask, condition_mask):
|
|
346
|
+
with torch.no_grad():
|
|
347
|
+
with torch.amp.autocast(self._device_type):
|
|
348
|
+
prediction, prediction_noise, latents_noise = self._model.forward(
|
|
349
|
+
input_ids,
|
|
350
|
+
conditions,
|
|
351
|
+
token_mask,
|
|
352
|
+
condition_mask,
|
|
353
|
+
)
|
|
354
|
+
loss = self._loss_fn(
|
|
355
|
+
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
356
|
+
)
|
|
357
|
+
return loss.detach()
|
|
358
|
+
|
|
359
|
+
def _save_state(self, checkpoint: str, epoch: int):
|
|
360
|
+
if _is_local():
|
|
361
|
+
os.makedirs(checkpoint, exist_ok=True)
|
|
362
|
+
self._model.save_pretrained(checkpoint) # type: ignore
|
|
363
|
+
torch.save(
|
|
364
|
+
{
|
|
365
|
+
"next_epoch": epoch + 1,
|
|
366
|
+
"optimizer": self._optimizer.state_dict(),
|
|
367
|
+
"scheduler": self._scheduler.state_dict(),
|
|
368
|
+
},
|
|
369
|
+
os.path.join(checkpoint, "train_state.pt"),
|
|
370
|
+
)
|
|
371
|
+
self.config_typed.to_json(os.path.join(checkpoint, "train_config.json"))
|
|
372
|
+
|
|
373
|
+
def _get_backend(self) -> tuple[str, str]:
|
|
374
|
+
if torch.accelerator.is_available():
|
|
375
|
+
device_type = torch.accelerator.current_accelerator().type # type: ignore
|
|
376
|
+
vendor_backend = torch.distributed.get_default_backend_for_device(
|
|
377
|
+
device_type
|
|
378
|
+
)
|
|
379
|
+
|
|
380
|
+
else:
|
|
381
|
+
device_type = torch.device("cpu").type
|
|
382
|
+
vendor_backend = torch.distributed.get_default_backend_for_device(
|
|
383
|
+
device_type
|
|
384
|
+
)
|
|
385
|
+
|
|
386
|
+
return device_type, vendor_backend
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: kernel_elastic_autoencoder
|
|
3
|
+
Version: 1.0.0
|
|
4
|
+
Summary: Implementation of Kernel-Elastic Autoencoder for Molecular Design (https://doi.org/10.1093/pnasnexus/pgae168)
|
|
5
|
+
License: MIT
|
|
6
|
+
Author: Felix Rotter-McCartney
|
|
7
|
+
Author-email: felix.rotter@mail.utoronto.ca
|
|
8
|
+
Requires-Python: >=3.12,<3.15
|
|
9
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
10
|
+
Classifier: Programming Language :: Python :: 3
|
|
11
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
12
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.14
|
|
14
|
+
Requires-Dist: huggingface-hub (>=1.22.0,<2.0.0)
|
|
15
|
+
Requires-Dist: pandas (>=3.0.3,<4.0.0)
|
|
16
|
+
Requires-Dist: pydantic (>=2.13.4,<3.0.0)
|
|
17
|
+
Requires-Dist: safetensors (>=0.8.0,<0.9.0)
|
|
18
|
+
Description-Content-Type: text/markdown
|
|
19
|
+
|
|
20
|
+
## `kernel_elastic_autoencoder`
|
|
21
|
+
|
|
22
|
+
`kernel_elastic_autoencoder` is a library implementing the architecture and techniques described
|
|
23
|
+
in [this publication](https://doi.org/10.1093/pnasnexus/pgae168) from Li et al. I am not affiliated with the authors of
|
|
24
|
+
the original paper, and this implementation is provided as-is, with no guarantee of completeness.
|
|
25
|
+
|
|
26
|
+
### Installation
|
|
27
|
+
|
|
28
|
+
`kernel_elastic_autoencoder` can be installed with pip, and regular builds are provided on PyPI:
|
|
29
|
+
|
|
30
|
+
pip install kernel_elastic_autoencoder
|
|
31
|
+
|
|
32
|
+
> Please note that `torch` is not included as a dependency due to its many hardware-accelerator-dependent versions, so
|
|
33
|
+
> take care to install the appropriate version manually.
|
|
34
|
+
|
|
35
|
+
Distribution builds are also provided here on GitHub Releases. New builds are triggered by the CD Action, so they will
|
|
36
|
+
be made available as soon as a new PR is merged to `main`.
|
|
37
|
+
|
|
38
|
+
Alternatively, for development purposes, `kernel_elastic_autoencoder` may be installed from source provided here. Builds
|
|
39
|
+
and deps are managed with `poetry`.
|
|
40
|
+
|
|
41
|
+
### Documentation
|
|
42
|
+
|
|
43
|
+
API documentation is generated with `pdoc` and covers the `__all__`-exported interfaces. It is available on
|
|
44
|
+
[GitHub Pages](https://cancelradius.github.io/kernel_elastic_autoencoder).
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
kernel_elastic_autoencoder/__init__.py,sha256=z7Vo43n0hcgYHHmPv2VTq9wvmRgD475F2Oa5cdysytQ,1227
|
|
2
|
+
kernel_elastic_autoencoder/collate.py,sha256=k0p5Eryr7g7HplK6V8Ur9Gdg1sUR298zFWd8F-hGhzM,6114
|
|
3
|
+
kernel_elastic_autoencoder/config.py,sha256=pm0i40_SS57Ur2cUpmJ1VoM20q-HY2D2veG8rULITm4,9052
|
|
4
|
+
kernel_elastic_autoencoder/layers.py,sha256=QvZkyxij5x674E4YyhDDKIzsTAmT8-URlApWW2Vm2UU,6823
|
|
5
|
+
kernel_elastic_autoencoder/losses.py,sha256=pMRKg-OmnuixLEz1yvmQw-EhYIGHU3wvVxGte4KY1Yo,6212
|
|
6
|
+
kernel_elastic_autoencoder/model.py,sha256=30kLX1tPhaEycF0UoCsqE9EABSQ7_6aAwQO_U4Z2uTg,9128
|
|
7
|
+
kernel_elastic_autoencoder/pipeline.py,sha256=39F294yrdBYPObPQp0r-rx2FlgPGEFZz8KiyHm_BUKI,8529
|
|
8
|
+
kernel_elastic_autoencoder/sample.py,sha256=kRx4fZzQ169vdjWbtY-21yoNNX_3YmegIKxWYFhPU8U,1802
|
|
9
|
+
kernel_elastic_autoencoder/tokenizer.py,sha256=-Fx-pV8s5rDxsZMbDKxwhz4Xmu8xUhIyAPtCE1Gii_E,4780
|
|
10
|
+
kernel_elastic_autoencoder/training.py,sha256=FDpk059D94kYkfsi5l-1yzTjHI5U8Hm6PAP1FKTfGVQ,13742
|
|
11
|
+
kernel_elastic_autoencoder-1.0.0.dist-info/METADATA,sha256=wkNaUu-lbFlxOmVqHfeabTCrfMg0iLr7QWnOh7tx_lY,2001
|
|
12
|
+
kernel_elastic_autoencoder-1.0.0.dist-info/WHEEL,sha256=eY7nduwzv-ldUxpzbRlxwvC693Hg6PX8bWDjEHjZ_dk,88
|
|
13
|
+
kernel_elastic_autoencoder-1.0.0.dist-info/RECORD,,
|