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,127 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import torch.nn.functional as F
|
|
3
|
+
from torch import nn
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class Loss(nn.Module):
|
|
7
|
+
"""Module for Kernel-Elastic Autoencoder loss.
|
|
8
|
+
|
|
9
|
+
Layer storing hyperparameters, computing loss on forward pass. Losses are computed as:
|
|
10
|
+
|
|
11
|
+
$\\mathcal{L}(\\lambda, \\delta)=\\mathcal{L}_{WCEL}(\\lambda, \\delta) + m$-$MMD(\\lambda)$
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
def __init__(
|
|
15
|
+
self,
|
|
16
|
+
*,
|
|
17
|
+
hp_lambda: float,
|
|
18
|
+
hp_delta: float,
|
|
19
|
+
hp_sigma: float,
|
|
20
|
+
kernel_dist_size: int,
|
|
21
|
+
padding_idx: int,
|
|
22
|
+
embedding_dim: int,
|
|
23
|
+
pooling_dim: int,
|
|
24
|
+
):
|
|
25
|
+
"""Instantiates a Loss object.
|
|
26
|
+
|
|
27
|
+
Args:
|
|
28
|
+
hp_lambda: Hyperparameter $\\lambda$, as used in WCEL and m-MMD losses. Roughly, controls how strongly the
|
|
29
|
+
shape of the latent vector distribution is penalized.
|
|
30
|
+
hp_delta: Hyperparameter $\\delta$, as used in WCEL loss. Roughly, controls the relative weights of the
|
|
31
|
+
vanilla-AE and VAE objectives in the reconstruction loss.
|
|
32
|
+
hp_sigma: Hyperparameter $\\sigma$, as used in the Kernel function applied in m-MMD loss. Roughly,
|
|
33
|
+
used as a scaling factor to control the sizes of gradients produced by the m-MMD loss.
|
|
34
|
+
kernel_dist_size: Size of the sampled distribution of vectors used to penalize the shape of the latent
|
|
35
|
+
vector distribution through the kernel.
|
|
36
|
+
padding_idx: Index of padding token. Used internally to zero vectors corresponding to padding tokens through
|
|
37
|
+
embedding layers. Should be fetched with Tokenizer.pad_token_id as specified in the Tokenizer protocol.
|
|
38
|
+
embedding_dim: Embedding dimension used by nn.Embedding layers.
|
|
39
|
+
pooling_dim: Sequence length dimension to which inputs are pooled after condition concatenation through the
|
|
40
|
+
encoder. Proportional to the dimension of latent vectors.
|
|
41
|
+
"""
|
|
42
|
+
super().__init__()
|
|
43
|
+
self.hp_lambda = hp_lambda
|
|
44
|
+
"""Hyperparameter $\\lambda$, as used in WCEL and m-MMD losses. Roughly, controls how strongly the
|
|
45
|
+
shape of the latent vector distribution is penalized."""
|
|
46
|
+
self.hp_delta = hp_delta
|
|
47
|
+
"""Hyperparameter $\\delta$, as used in WCEL loss. Roughly, controls the relative weights of the
|
|
48
|
+
vanilla-AE and VAE objectives in the reconstruction loss."""
|
|
49
|
+
self.hp_sigma = hp_sigma
|
|
50
|
+
"""Hyperparameter $\\sigma$, as used in the Kernel function applied in m-MMD loss. Roughly,
|
|
51
|
+
used as a scaling factor to control the sizes of gradients produced by the m-MMD loss."""
|
|
52
|
+
self.kernel_dist_size = kernel_dist_size
|
|
53
|
+
"""Size of the sampled distribution of vectors used to penalize the shape of the latent
|
|
54
|
+
vector distribution through the kernel."""
|
|
55
|
+
self.padding_idx = padding_idx
|
|
56
|
+
"""Index of padding token. Used internally to zero vectors corresponding to padding tokens through
|
|
57
|
+
embedding layers. Should be fetched with Tokenizer.pad_token_id as specified in the Tokenizer protocol."""
|
|
58
|
+
self.embedding_dim = embedding_dim
|
|
59
|
+
"""Embedding dimension used by nn.Embedding layers."""
|
|
60
|
+
self.pooling_dim = pooling_dim
|
|
61
|
+
"""Sequence length dimension to which inputs are pooled after condition concatenation through the
|
|
62
|
+
encoder. Proportional to the dimension of latent vectors."""
|
|
63
|
+
|
|
64
|
+
self.register_buffer("_loc", torch.zeros(self.pooling_dim * self.embedding_dim))
|
|
65
|
+
self.register_buffer("_cov", torch.eye(self.pooling_dim * self.embedding_dim))
|
|
66
|
+
|
|
67
|
+
def forward(
|
|
68
|
+
self,
|
|
69
|
+
prediction: torch.Tensor,
|
|
70
|
+
prediction_noise: torch.Tensor,
|
|
71
|
+
ground_truth: torch.Tensor,
|
|
72
|
+
latents: torch.Tensor,
|
|
73
|
+
) -> torch.Tensor:
|
|
74
|
+
"""Forward pass through the loss module. May also be called using Loss.__call__.
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
prediction: Tensor of dimension (B, M, L) containing teacher forcing prediction logits without noise
|
|
78
|
+
before the decoder.
|
|
79
|
+
prediction_noise: Tensor of dimension (B, M, L) containing teacher forcing prediction logits with added noise
|
|
80
|
+
before the decoder.
|
|
81
|
+
ground_truth: Tensor of dimension (B, M) containing the indices of ground truth tokens for each sequence.
|
|
82
|
+
latents: Tensor of dimension (B, P * E) containing latent vectors produced by the encoder with added
|
|
83
|
+
noise.
|
|
84
|
+
|
|
85
|
+
Returns:
|
|
86
|
+
torch.Tensor: Tensor of dimension (1) containing loss value.
|
|
87
|
+
"""
|
|
88
|
+
return self._weighted_cross_entropy_loss(
|
|
89
|
+
prediction, prediction_noise, ground_truth
|
|
90
|
+
) + self._modified_maximum_mean_discrepancy(latents)
|
|
91
|
+
|
|
92
|
+
def _weighted_cross_entropy_loss(
|
|
93
|
+
self,
|
|
94
|
+
prediction: torch.Tensor,
|
|
95
|
+
prediction_noise: torch.Tensor,
|
|
96
|
+
ground_truth: torch.Tensor,
|
|
97
|
+
) -> torch.Tensor:
|
|
98
|
+
ground_truth.to(torch.long)
|
|
99
|
+
log_softmax = torch.log_softmax(prediction, dim=-1)
|
|
100
|
+
log_softmax_noise = torch.log_softmax(prediction_noise, dim=-1)
|
|
101
|
+
Y = F.one_hot(ground_truth, num_classes=prediction.size(-1))
|
|
102
|
+
Y[..., self.padding_idx] = 0
|
|
103
|
+
t1 = (Y * log_softmax_noise).sum(dim=(-1, -2))
|
|
104
|
+
t2 = (Y * log_softmax).sum(dim=(-1, -2))
|
|
105
|
+
coeff = -1 / (self.hp_lambda + self.hp_delta + 1)
|
|
106
|
+
loss = coeff * (t1 + ((self.hp_lambda + self.hp_delta) * t2))
|
|
107
|
+
return loss.mean()
|
|
108
|
+
|
|
109
|
+
def _modified_maximum_mean_discrepancy(
|
|
110
|
+
self,
|
|
111
|
+
latents: torch.Tensor,
|
|
112
|
+
) -> torch.Tensor:
|
|
113
|
+
mvn = torch.distributions.MultivariateNormal(
|
|
114
|
+
loc=self._loc, # type: ignore
|
|
115
|
+
covariance_matrix=self._cov, # type: ignore
|
|
116
|
+
)
|
|
117
|
+
samples = mvn.rsample((self.kernel_dist_size,))
|
|
118
|
+
square_difference_sum = torch.cdist(latents, samples, p=2.0).pow(2)
|
|
119
|
+
kernel_pairwise_sum = torch.exp(
|
|
120
|
+
((-1 / (self.pooling_dim * self.embedding_dim)) * square_difference_sum)
|
|
121
|
+
/ (2 * (self.hp_sigma**2))
|
|
122
|
+
).sum()
|
|
123
|
+
loss = self.hp_lambda * (
|
|
124
|
+
torch.tensor([1])
|
|
125
|
+
- ((1 / (latents.size(0) * self.kernel_dist_size)) * kernel_pairwise_sum)
|
|
126
|
+
)
|
|
127
|
+
return loss
|
|
@@ -0,0 +1,222 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
from typing import Any
|
|
3
|
+
|
|
4
|
+
import torch
|
|
5
|
+
from huggingface_hub import PyTorchModelHubMixin
|
|
6
|
+
from huggingface_hub.hub_mixin import T, DataclassInstance
|
|
7
|
+
from torch import nn
|
|
8
|
+
|
|
9
|
+
from kernel_elastic_autoencoder.config import ModelConfig
|
|
10
|
+
from kernel_elastic_autoencoder.layers import Decoder, Encoder, TrainingNoise
|
|
11
|
+
|
|
12
|
+
CODERS = {
|
|
13
|
+
ModelConfig: (
|
|
14
|
+
lambda x: x.model_dump(),
|
|
15
|
+
lambda data: ModelConfig.model_validate(data),
|
|
16
|
+
)
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Model(
|
|
21
|
+
nn.Module,
|
|
22
|
+
PyTorchModelHubMixin,
|
|
23
|
+
coders=CODERS,
|
|
24
|
+
library_name="kernel_elastic_autoencoder",
|
|
25
|
+
repo_url="https://github.com/cancelradius/kae",
|
|
26
|
+
paper_url="https://arxiv.org/abs/2310.08685",
|
|
27
|
+
):
|
|
28
|
+
"""Base class for Kernel-Elastic Autoencoder models.
|
|
29
|
+
|
|
30
|
+
Defines a barebones API for model calls with tensor inputs. For easy use with standard data formats, calling Model
|
|
31
|
+
through the Pipeline API may be preferred.
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
def __init__(self, config: dict | ModelConfig):
|
|
35
|
+
"""Instantiates a Model object.
|
|
36
|
+
|
|
37
|
+
Args:
|
|
38
|
+
config: Dictionary or ModelConfig schema defining model parameters. Will be validated with ModelConfig
|
|
39
|
+
regardless of input type.
|
|
40
|
+
"""
|
|
41
|
+
super().__init__()
|
|
42
|
+
|
|
43
|
+
self.config = config
|
|
44
|
+
"""Configuration object for Hugging Face Hub compatible serialization. Not recommended to use, as it is
|
|
45
|
+
internal-use. Use Model.config_typed instead."""
|
|
46
|
+
config_typed = ModelConfig.model_validate(config)
|
|
47
|
+
self.config_typed: ModelConfig = config_typed
|
|
48
|
+
"""Type-validated config in a Pydantic ModelConfig schema, recommended for public API use."""
|
|
49
|
+
|
|
50
|
+
self.encoder = Encoder(
|
|
51
|
+
max_len=config_typed.input.max_len,
|
|
52
|
+
vocab_size=config_typed.input.vocab_size,
|
|
53
|
+
condition_channels=config_typed.input.condition_channels,
|
|
54
|
+
embedding_dim=config_typed.common.embedding_dim,
|
|
55
|
+
pooling_dim=config_typed.common.pooling_dim,
|
|
56
|
+
padding_idx=config_typed.common.padding_idx,
|
|
57
|
+
padding_value=config_typed.common.padding_value,
|
|
58
|
+
num_layers=config_typed.encoder.num_layers,
|
|
59
|
+
num_heads=config_typed.encoder.num_heads,
|
|
60
|
+
feedforward_scale=config_typed.encoder.feedforward_scale,
|
|
61
|
+
dropout=config_typed.encoder.dropout,
|
|
62
|
+
)
|
|
63
|
+
"""Model encoder module."""
|
|
64
|
+
self.noise = TrainingNoise()
|
|
65
|
+
"""Noise-adding module for training purposes."""
|
|
66
|
+
self.decoder = Decoder(
|
|
67
|
+
max_len=config_typed.input.max_len,
|
|
68
|
+
vocab_size=config_typed.input.vocab_size,
|
|
69
|
+
condition_channels=config_typed.input.condition_channels,
|
|
70
|
+
embedding_dim=config_typed.common.embedding_dim,
|
|
71
|
+
pooling_dim=config_typed.common.pooling_dim,
|
|
72
|
+
padding_idx=config_typed.common.padding_idx,
|
|
73
|
+
padding_value=config_typed.common.padding_value,
|
|
74
|
+
num_layers=config_typed.decoder.num_layers,
|
|
75
|
+
num_heads=config_typed.decoder.num_heads,
|
|
76
|
+
feedforward_scale=config_typed.decoder.feedforward_scale,
|
|
77
|
+
dropout=config_typed.decoder.dropout,
|
|
78
|
+
)
|
|
79
|
+
"""Model decoder module."""
|
|
80
|
+
|
|
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
|
+
def forward(
|
|
123
|
+
self,
|
|
124
|
+
input_ids: torch.Tensor,
|
|
125
|
+
conditions: torch.Tensor,
|
|
126
|
+
token_mask: torch.Tensor,
|
|
127
|
+
condition_mask: torch.Tensor,
|
|
128
|
+
):
|
|
129
|
+
"""Forward pass for training loop. Produces all tensors needed for computing loss with
|
|
130
|
+
Loss.forward.
|
|
131
|
+
|
|
132
|
+
Args:
|
|
133
|
+
input_ids: Tensor of dimension (B, M) containing input ids for each sequence.
|
|
134
|
+
conditions: Tensor of dimension (B, C) containing condition values for each sequence.
|
|
135
|
+
token_mask: Tensor of dimension (B, M) containing boolean padding masks for each sequence.
|
|
136
|
+
condition_mask: Tensor of dimension (B, C) containing boolean condition padding masks for each sequence.
|
|
137
|
+
|
|
138
|
+
Returns:
|
|
139
|
+
torch.Tensor: Tensor of dimension (B, M, L) containing teacher forcing prediction logits without noise
|
|
140
|
+
before the decoder.
|
|
141
|
+
torch.Tensor: Tensor of dimension (B, M, L) containing teacher forcing prediction logits with added noise
|
|
142
|
+
before the decoder.
|
|
143
|
+
torch.Tensor: Tensor of dimension (B, P * E) containing latent vectors produced by the encoder with added
|
|
144
|
+
noise.
|
|
145
|
+
"""
|
|
146
|
+
padding_mask = torch.cat((token_mask, condition_mask), 1)
|
|
147
|
+
latents, condition_embeddings = self.encoder(
|
|
148
|
+
input_ids, conditions, padding_mask
|
|
149
|
+
)
|
|
150
|
+
latents_noise = self.noise(latents)
|
|
151
|
+
|
|
152
|
+
prediction = self.decoder(
|
|
153
|
+
input_ids[:, :-1],
|
|
154
|
+
latents,
|
|
155
|
+
condition_embeddings,
|
|
156
|
+
padding_mask[: self.config_typed.input.max_len - 1],
|
|
157
|
+
)
|
|
158
|
+
prediction_noise = self.decoder(
|
|
159
|
+
input_ids[:, :-1],
|
|
160
|
+
latents_noise,
|
|
161
|
+
condition_embeddings,
|
|
162
|
+
padding_mask[: self.config_typed.input.max_len - 1],
|
|
163
|
+
)
|
|
164
|
+
return prediction, prediction_noise, latents_noise
|
|
165
|
+
|
|
166
|
+
def encode(
|
|
167
|
+
self,
|
|
168
|
+
input_ids: torch.Tensor,
|
|
169
|
+
conditions: torch.Tensor,
|
|
170
|
+
token_mask: torch.Tensor | None,
|
|
171
|
+
condition_mask: torch.Tensor,
|
|
172
|
+
) -> torch.Tensor:
|
|
173
|
+
"""Basic interface for a forward pass through the Model.model.encoder module.
|
|
174
|
+
|
|
175
|
+
Args:
|
|
176
|
+
input_ids: Tensor of dimension (B, M) containing input ids for each sequence.
|
|
177
|
+
conditions: Tensor of dimension (B, C) containing condition values for each sequence.
|
|
178
|
+
token_mask: Tensor of dimension (B, M) containing boolean padding masks for each sequence.
|
|
179
|
+
condition_mask: Tensor of dimension (B, C) containing boolean condition padding masks for each sequence.
|
|
180
|
+
|
|
181
|
+
Returns:
|
|
182
|
+
torch.Tensor: Tensor of dimension (B, P * E) containing latent vectors produced by the encoder.
|
|
183
|
+
"""
|
|
184
|
+
padding_mask = (
|
|
185
|
+
torch.cat((token_mask, condition_mask), 1)
|
|
186
|
+
if (token_mask is not None)
|
|
187
|
+
else torch.cat((torch.full_like(input_ids, True), condition_mask), 1)
|
|
188
|
+
)
|
|
189
|
+
return self.encoder(input_ids, conditions, padding_mask)[0]
|
|
190
|
+
|
|
191
|
+
def decode(
|
|
192
|
+
self,
|
|
193
|
+
current_output: torch.Tensor,
|
|
194
|
+
latents: torch.Tensor,
|
|
195
|
+
condition_embeddings: torch.Tensor,
|
|
196
|
+
token_mask: torch.Tensor | None,
|
|
197
|
+
condition_mask: torch.Tensor,
|
|
198
|
+
) -> torch.Tensor:
|
|
199
|
+
"""Basic interface for a forward pass through the Model.model.decoder module.
|
|
200
|
+
|
|
201
|
+
Args:
|
|
202
|
+
current_output: Tensor of dimension (B, S) containing previous output ids for each sequence. In
|
|
203
|
+
autoregressive generation, this should initially represent a single start-of-sequence token.
|
|
204
|
+
latents: Tensor of dimension (B, P * E) containing latent vectors. May be produced using Model.encode.
|
|
205
|
+
condition_embeddings: Tensor of dimension (B, C, E) containing condition embeddings for each sequence. May
|
|
206
|
+
be produced using Model.embed_conditions.
|
|
207
|
+
token_mask: Tensor of dimension (B, S) containing boolean padding masks for each sequence. If None is
|
|
208
|
+
passed, the absence of padding is assumed.
|
|
209
|
+
condition_mask: Tensor of dimension (B, C) containing boolean condition padding masks for each sequence.
|
|
210
|
+
|
|
211
|
+
Returns:
|
|
212
|
+
torch.Tensor: Tensor of dimension (B, S, L) containing prediction logits produced by the decoder.
|
|
213
|
+
"""
|
|
214
|
+
padding_mask = (
|
|
215
|
+
torch.cat((token_mask, condition_mask), 1)
|
|
216
|
+
if (token_mask is not None)
|
|
217
|
+
else torch.cat((torch.full_like(current_output, False), condition_mask), 1)
|
|
218
|
+
).to(torch.bool)
|
|
219
|
+
return self.decoder(current_output, latents, condition_embeddings, padding_mask)
|
|
220
|
+
|
|
221
|
+
def embed_conditions(self, conditions):
|
|
222
|
+
return self.encoder.embedding.conditional_embedding(conditions)
|
|
@@ -0,0 +1,203 @@
|
|
|
1
|
+
from collections.abc import Callable, Iterable
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
from pydantic import BaseModel, ConfigDict
|
|
5
|
+
|
|
6
|
+
from kernel_elastic_autoencoder.collate import Collator, DataframeCollator
|
|
7
|
+
from kernel_elastic_autoencoder.config import ModelConfig
|
|
8
|
+
from kernel_elastic_autoencoder.model import Model
|
|
9
|
+
from kernel_elastic_autoencoder.sample import Sampler, Top1Sampler
|
|
10
|
+
from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class CompletionIntermediate(BaseModel):
|
|
14
|
+
model_config = ConfigDict(arbitrary_types_allowed=True)
|
|
15
|
+
latents: torch.Tensor
|
|
16
|
+
input_ids: torch.Tensor
|
|
17
|
+
batches_completed: torch.Tensor
|
|
18
|
+
condition_embeddings: torch.Tensor
|
|
19
|
+
condition_mask: torch.Tensor
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class Completion(BaseModel):
|
|
23
|
+
"""Return schema for Pipeline.completion."""
|
|
24
|
+
|
|
25
|
+
model_config = ConfigDict(arbitrary_types_allowed=True)
|
|
26
|
+
outputs: Iterable[str]
|
|
27
|
+
"""Iterable of generated output sequences."""
|
|
28
|
+
condition_embeddings: torch.Tensor
|
|
29
|
+
"""Tensor containing generated condition embeddings for reuse."""
|
|
30
|
+
condition_mask: torch.Tensor
|
|
31
|
+
"""Tensor containing generated condition masks for reuse."""
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class Pipeline[_T]:
|
|
35
|
+
"""User-facing pipeline for inference.
|
|
36
|
+
|
|
37
|
+
Defines an easy-to-use API for decoder-only inference with a pretrained model.
|
|
38
|
+
|
|
39
|
+
Examples:
|
|
40
|
+
Wrapping a pretrained model and tokenizer:
|
|
41
|
+
>>> model = Model.from_pretrained("./checkpoint")
|
|
42
|
+
>>> tokenizer = MyClassImplementingTokenizer.from_pretrained("./checkpoint/tokenizer")
|
|
43
|
+
>>> pipe = Pipeline(model, tokenizer, MyClassImplementingCollator, MyClassImplementingSampler)
|
|
44
|
+
>>> compl = pipe.completion(latents, dataset, seq_feature, cond_features)
|
|
45
|
+
>>> print(compl.outputs)
|
|
46
|
+
|
|
47
|
+
TODO: Complete examples.
|
|
48
|
+
"""
|
|
49
|
+
|
|
50
|
+
def __init__(
|
|
51
|
+
self,
|
|
52
|
+
model: Model,
|
|
53
|
+
tokenizer: Tokenizer,
|
|
54
|
+
collator: Callable[[ModelConfig, Tokenizer], Collator[_T]] = DataframeCollator,
|
|
55
|
+
sampler: Callable[[Tokenizer], Sampler] = Top1Sampler,
|
|
56
|
+
device: torch.device | None = None,
|
|
57
|
+
) -> None:
|
|
58
|
+
"""Instantiates a Pipeline object.
|
|
59
|
+
|
|
60
|
+
TODO: Pipeline setup will be streamlined such that collators and samplers are by default fetched from model
|
|
61
|
+
configuration.
|
|
62
|
+
|
|
63
|
+
Args:
|
|
64
|
+
model: Pre-trained Model object for inference. Can be obtained from Model.from_pretrained.
|
|
65
|
+
tokenizer: Pre-configured Tokenizer object. The Tokenizer protocol supports tokenizers inheriting from
|
|
66
|
+
transformers.PreTrainedTokenizerBase, so such tokenizers may be loaded from HuggingFace Hub.
|
|
67
|
+
collator: Constructor for a data collator inheriting from Collator. Determines input types to Pipeline methods.
|
|
68
|
+
sampler: Constructor for a data sampler inheriting from Sampler. Determines how outputs are produced from
|
|
69
|
+
Pipeline methods.
|
|
70
|
+
device: Torch device used for inference.
|
|
71
|
+
"""
|
|
72
|
+
self.model = model.to(device)
|
|
73
|
+
"""Pre-trained Model object for inference. Moved to Pipeline.device, and placed in eval() mode."""
|
|
74
|
+
self.model.eval()
|
|
75
|
+
self.tokenizer = tokenizer
|
|
76
|
+
"""Pre-configured Tokenizer object. The Tokenizer Protocol supports tokenizers inheriting from
|
|
77
|
+
transformers.PreTrainedTokenizerBase, so such tokenizers may be loaded from HuggingFace Hub."""
|
|
78
|
+
self.device = device
|
|
79
|
+
"""Torch device used for inference."""
|
|
80
|
+
|
|
81
|
+
self.collator = collator(model.config_typed, tokenizer)
|
|
82
|
+
"""Data collator inheriting from Collator, instantiated by passing Pipeline.model.config_typed and Pipeline.tokenizer."""
|
|
83
|
+
self.sampler = sampler(tokenizer)
|
|
84
|
+
"""Data sampler inheriting from Sampler, instantiated by passing Pipeline.tokenizer."""
|
|
85
|
+
|
|
86
|
+
@torch.inference_mode()
|
|
87
|
+
def _completion_entry(
|
|
88
|
+
self,
|
|
89
|
+
latents: torch.Tensor,
|
|
90
|
+
dataset: _T,
|
|
91
|
+
seq_feature: str,
|
|
92
|
+
cond_features: list[str],
|
|
93
|
+
device: torch.device | None = None,
|
|
94
|
+
**kwargs,
|
|
95
|
+
) -> CompletionIntermediate:
|
|
96
|
+
ds = self.collator(
|
|
97
|
+
dataset=dataset,
|
|
98
|
+
seq_feature=seq_feature,
|
|
99
|
+
cond_features=cond_features,
|
|
100
|
+
padding=False,
|
|
101
|
+
add_special_tokens=False,
|
|
102
|
+
device=device,
|
|
103
|
+
**kwargs,
|
|
104
|
+
)
|
|
105
|
+
conds_embed = self.model.embed_conditions(ds.conditions)
|
|
106
|
+
return CompletionIntermediate(
|
|
107
|
+
latents=latents,
|
|
108
|
+
input_ids=ds.input_ids,
|
|
109
|
+
batches_completed=torch.zeros(ds.input_ids.size(0), dtype=torch.bool),
|
|
110
|
+
condition_embeddings=conds_embed,
|
|
111
|
+
condition_mask=ds.condition_mask,
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
@torch.inference_mode()
|
|
115
|
+
def _completion_step(
|
|
116
|
+
self, intermediate: CompletionIntermediate, **kwargs
|
|
117
|
+
) -> CompletionIntermediate:
|
|
118
|
+
intermediate.input_ids = torch.cat(
|
|
119
|
+
[
|
|
120
|
+
torch.full(
|
|
121
|
+
(intermediate.input_ids.size(0), 1), self.tokenizer.bos_token_id
|
|
122
|
+
),
|
|
123
|
+
intermediate.input_ids,
|
|
124
|
+
],
|
|
125
|
+
dim=1,
|
|
126
|
+
)
|
|
127
|
+
logits = self.model.decode(
|
|
128
|
+
current_output=intermediate.input_ids,
|
|
129
|
+
latents=intermediate.latents,
|
|
130
|
+
condition_embeddings=intermediate.condition_embeddings,
|
|
131
|
+
token_mask=None,
|
|
132
|
+
condition_mask=intermediate.condition_mask,
|
|
133
|
+
)
|
|
134
|
+
new_toks = self.sampler.sample_ids(logits[:, -1:])
|
|
135
|
+
intermediate.batches_completed |= (
|
|
136
|
+
new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
137
|
+
)
|
|
138
|
+
torch.where(
|
|
139
|
+
intermediate.batches_completed.unsqueeze(-1),
|
|
140
|
+
self.tokenizer.pad_token_id,
|
|
141
|
+
new_toks,
|
|
142
|
+
)
|
|
143
|
+
intermediate.input_ids = torch.cat([intermediate.input_ids, new_toks], dim=1)
|
|
144
|
+
return intermediate
|
|
145
|
+
|
|
146
|
+
@torch.inference_mode()
|
|
147
|
+
def _completion_exit(
|
|
148
|
+
self, intermediate: CompletionIntermediate, **kwargs
|
|
149
|
+
) -> Completion:
|
|
150
|
+
return Completion(
|
|
151
|
+
outputs=self.tokenizer.decode(
|
|
152
|
+
intermediate.input_ids, skip_special_tokens=True
|
|
153
|
+
),
|
|
154
|
+
condition_embeddings=intermediate.condition_embeddings,
|
|
155
|
+
condition_mask=intermediate.condition_mask,
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
def completion(
|
|
159
|
+
self,
|
|
160
|
+
latents: torch.Tensor,
|
|
161
|
+
dataset: _T,
|
|
162
|
+
seq_feature: str,
|
|
163
|
+
cond_features: list[str],
|
|
164
|
+
device: torch.device | None = None,
|
|
165
|
+
**kwargs,
|
|
166
|
+
) -> Completion:
|
|
167
|
+
"""Completes each batch in an input dataset with latent and condition guidance through the decoder.
|
|
168
|
+
|
|
169
|
+
Using decoder-only inference, the model completes each batch in the provided dataset. Dataset input typing is
|
|
170
|
+
uniquely defined by the used Collator. The input dataset is passed through the collator, then subject to
|
|
171
|
+
autoregressive decoder-only inference through the model, where new tokens are sampled from intermediates and
|
|
172
|
+
stored in a tensor format. Finally, on reaching end-of-sequence tokens or the maximum sequence length, the
|
|
173
|
+
sequence is decoded by the tokenizer and returned along with the generated condition embeddings and masks in a
|
|
174
|
+
Completion schema.
|
|
175
|
+
|
|
176
|
+
TODO: Callbacks are planned to enable easy custom inference functions and streaming of text.
|
|
177
|
+
|
|
178
|
+
Args:
|
|
179
|
+
latents: Tensor of dimension (B, P * E) containing latent vectors for the batches.
|
|
180
|
+
dataset: Dataset of type specified in Pipeline.collator.
|
|
181
|
+
seq_feature: String to pass to collator containing the feature or column containing text sequences to
|
|
182
|
+
complete.
|
|
183
|
+
cond_features: List of strings to pass to collator containing features or columns containing numerical
|
|
184
|
+
condition values.
|
|
185
|
+
device: Torch device used for inference.
|
|
186
|
+
**kwargs: Additional keyword arguments passed to Pipeline.collator.
|
|
187
|
+
|
|
188
|
+
Returns:
|
|
189
|
+
Completion: Completion object, containing outputs, as well as condition embeddings and masks for reuse.
|
|
190
|
+
"""
|
|
191
|
+
intermediate = self._completion_entry(
|
|
192
|
+
latents=latents,
|
|
193
|
+
dataset=dataset,
|
|
194
|
+
seq_feature=seq_feature,
|
|
195
|
+
cond_features=cond_features,
|
|
196
|
+
device=device,
|
|
197
|
+
**kwargs,
|
|
198
|
+
)
|
|
199
|
+
while (
|
|
200
|
+
intermediate.input_ids.size(1) < self.model.config_typed.input.max_len
|
|
201
|
+
) and (not intermediate.batches_completed.all()):
|
|
202
|
+
intermediate = self._completion_step(intermediate=intermediate)
|
|
203
|
+
return self._completion_exit(intermediate=intermediate)
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from collections.abc import Iterable
|
|
3
|
+
|
|
4
|
+
import torch
|
|
5
|
+
|
|
6
|
+
from kernel_elastic_autoencoder.tokenizer import Tokenizer
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class Sampler(ABC):
|
|
10
|
+
"""Base class for samplers.
|
|
11
|
+
|
|
12
|
+
Provides a specification for sampling text sequences from logits.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
def __init__(self, tokenizer: Tokenizer):
|
|
16
|
+
"""Instantiates a Sampler object.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
tokenizer: Object implementing the Tokenizer protocol.
|
|
20
|
+
"""
|
|
21
|
+
self.tokenizer = tokenizer
|
|
22
|
+
"""Object implementing the Tokenizer protocol."""
|
|
23
|
+
|
|
24
|
+
def __call__(
|
|
25
|
+
self, logits: torch.Tensor, skip_special_tokens: bool, **kwargs
|
|
26
|
+
) -> Iterable[str]:
|
|
27
|
+
ids = self.sample_ids(logits, **kwargs)
|
|
28
|
+
sampled = self.tokenizer.decode(ids, skip_special_tokens=skip_special_tokens)
|
|
29
|
+
return sampled
|
|
30
|
+
|
|
31
|
+
@abstractmethod
|
|
32
|
+
def sample_ids(self, logits: torch.Tensor, **kwargs) -> torch.Tensor:
|
|
33
|
+
"""Interface method for implementing index sampling from logits.
|
|
34
|
+
|
|
35
|
+
Args:
|
|
36
|
+
logits: Tensor of dimension (B, S, L) containing logits.
|
|
37
|
+
**kwargs: Keyword arguments.
|
|
38
|
+
|
|
39
|
+
Returns:
|
|
40
|
+
torch.Tensor: Tensor of dimension (B, S) containing vocabulary indices.
|
|
41
|
+
"""
|
|
42
|
+
...
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class Top1Sampler(Sampler):
|
|
46
|
+
def sample_ids(self, logits: torch.Tensor, **kwargs) -> torch.Tensor:
|
|
47
|
+
"""Implementation of index sampling from logits choosing the highest-probability token.
|
|
48
|
+
|
|
49
|
+
Args:
|
|
50
|
+
logits: Tensor of dimension (B, S, L) containing logits.
|
|
51
|
+
**kwargs: Keyword arguments.
|
|
52
|
+
|
|
53
|
+
Returns:
|
|
54
|
+
torch.Tensor: Tensor of dimension (B, S) containing vocabulary indices.
|
|
55
|
+
"""
|
|
56
|
+
return (
|
|
57
|
+
torch.topk(logits, k=1, dim=-1, **kwargs).indices.squeeze(-1).to(torch.long)
|
|
58
|
+
)
|