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.
@@ -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
+ )