kernel-elastic-autoencoder 3.0.0__tar.gz → 3.1.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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 3.0.0
3
+ Version: 3.1.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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "kernel_elastic_autoencoder"
3
- version = "3.0.0"
3
+ version = "3.1.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" }
@@ -0,0 +1,247 @@
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)
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 completion(
49
+ self,
50
+ latents: torch.Tensor,
51
+ sequences: list[str],
52
+ conditions: list[list[float]] | torch.Tensor,
53
+ device: torch.device | None = None,
54
+ **kwargs,
55
+ ) -> Iterable[str]:
56
+ """Completes each conditioned input sequence.
57
+
58
+ The model completes each sequence in the provided list using decoder-only inference. Tokens are
59
+ sampled greedily.
60
+
61
+ Args:
62
+ latents: Tensor of dimension (B, P * E) containing latent vectors for the batch.
63
+ sequences: List of text sequences to complete.
64
+ conditions: List of condition value lists per batch.
65
+ device: Torch device used for inference.
66
+ **kwargs: Additional keyword arguments passed to Tokenizer.encode.
67
+
68
+ Returns:
69
+ Iterable[str]: List of completed sequences, stripped of special tokens.
70
+ """
71
+ input_ids = self.tokenizer.encode(
72
+ seq=sequences,
73
+ padding=False,
74
+ max_length=self.model.config_typed.input.max_len,
75
+ add_special_tokens=False,
76
+ **kwargs,
77
+ ).to(self.device)
78
+ conditions = torch.as_tensor(conditions, dtype=torch.float, device=device)
79
+ condition_mask = (
80
+ conditions != self.model.config_typed.common.padding_value
81
+ ).to(torch.bool)
82
+ conds_embed = self.model.embed_conditions(conditions)
83
+ batches_completed = torch.zeros(input_ids.size(0), dtype=torch.bool)
84
+
85
+ while (input_ids.size(1) < self.model.config_typed.input.max_len) and (
86
+ not batches_completed.all()
87
+ ):
88
+ input_ids = torch.cat(
89
+ [
90
+ torch.full(
91
+ (input_ids.size(0), 1),
92
+ self.tokenizer.bos_token_id,
93
+ ),
94
+ input_ids,
95
+ ],
96
+ dim=1,
97
+ )
98
+ logits = self.model.decode(
99
+ current_output=input_ids,
100
+ latents=latents,
101
+ condition_embeddings=conds_embed,
102
+ token_mask=None,
103
+ condition_mask=condition_mask,
104
+ )
105
+ new_toks = (
106
+ torch.topk(logits[:, -1:], k=1, dim=-1)
107
+ .indices.squeeze(-1)
108
+ .to(torch.long)
109
+ )
110
+ batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
111
+ new_toks = torch.where(
112
+ batches_completed.unsqueeze(-1),
113
+ self.tokenizer.pad_token_id,
114
+ new_toks,
115
+ )
116
+ input_ids = torch.cat([input_ids, new_toks], dim=1)
117
+ return self.tokenizer.decode(input_ids, skip_special_tokens=True)
118
+
119
+ def beam_completion(
120
+ self,
121
+ latents: torch.Tensor,
122
+ beam_size: int,
123
+ sequences: list[str],
124
+ conditions: list[list[float]] | torch.Tensor,
125
+ device: torch.device | None = None,
126
+ **kwargs,
127
+ ) -> Iterable[str]:
128
+ """Completes each conditioned input sequence, using the beam search strategy.
129
+
130
+ The model completes each sequence in the provided list using decoder-only inference. On the
131
+ first step, the top `beam_size` first tokens are chosen for each batch. Subsequent tokens are
132
+ sampled greedily in parallel for every subsequence. Before returning output, the subsequence with
133
+ the highest sum of token logits is chosen for each batch.
134
+
135
+ Args:
136
+ latents: Tensor of dimension (B, P * E) containing latent vectors for the batch.
137
+ beam_size: Beam size of first step.
138
+ sequences: List of text sequences to complete.
139
+ conditions: List of condition value lists per batch.
140
+ device: Torch device used for inference.
141
+ **kwargs: Additional keyword arguments passed to Tokenizer.encode.
142
+
143
+ Returns:
144
+ Iterable[str]: List of completed sequences, stripped of special tokens.
145
+ """
146
+ input_ids = self.tokenizer.encode(
147
+ seq=sequences,
148
+ padding=False,
149
+ max_length=self.model.config_typed.input.max_len,
150
+ add_special_tokens=False,
151
+ **kwargs,
152
+ ).to(self.device)
153
+ input_probs = torch.zeros_like(input_ids * beam_size)
154
+ conditions = torch.as_tensor(conditions, dtype=torch.float, device=device)
155
+ condition_mask = (
156
+ (conditions != self.model.config_typed.common.padding_value)
157
+ .to(torch.bool)
158
+ .repeat_interleave(beam_size, dim=0)
159
+ )
160
+ conds_embed = self.model.embed_conditions(conditions).repeat_interleave(
161
+ beam_size, dim=0
162
+ )
163
+ batches_completed = torch.zeros(input_ids.size(0) * beam_size, dtype=torch.bool)
164
+
165
+ logits = self.model.decode(
166
+ current_output=input_ids,
167
+ latents=latents,
168
+ condition_embeddings=conds_embed,
169
+ token_mask=None,
170
+ condition_mask=condition_mask,
171
+ )
172
+ new_toks = (
173
+ torch.topk(logits[:, -1:], k=beam_size, dim=-1)
174
+ .indices.squeeze(-1)
175
+ .to(torch.long)
176
+ )
177
+ new_probs = (
178
+ torch.topk(logits[:, -1:], k=beam_size, dim=-1)
179
+ .values.squeeze(-1)
180
+ .to(torch.long)
181
+ )
182
+ batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
183
+ new_toks = torch.where(
184
+ batches_completed.unsqueeze(-1),
185
+ self.tokenizer.pad_token_id,
186
+ new_toks,
187
+ )
188
+ new_probs = torch.where(
189
+ batches_completed.unsqueeze(-1),
190
+ 0.0,
191
+ new_probs,
192
+ )
193
+ input_ids = torch.cat(
194
+ [input_ids.repeat_interleave(beam_size, dim=0), new_toks], dim=1
195
+ )
196
+ input_probs = torch.cat([input_probs, new_probs], dim=1)
197
+
198
+ while (input_ids.size(1) < self.model.config_typed.input.max_len) and (
199
+ not batches_completed.all()
200
+ ):
201
+ input_ids = torch.cat(
202
+ [
203
+ torch.full(
204
+ (input_ids.size(0), 1),
205
+ self.tokenizer.bos_token_id,
206
+ ),
207
+ input_ids,
208
+ ],
209
+ dim=1,
210
+ )
211
+ logits = self.model.decode(
212
+ current_output=input_ids,
213
+ latents=latents,
214
+ condition_embeddings=conds_embed,
215
+ token_mask=None,
216
+ condition_mask=condition_mask,
217
+ )
218
+ new_toks = (
219
+ torch.topk(logits[:, -1:], k=1, dim=-1)
220
+ .indices.squeeze(-1)
221
+ .to(torch.long)
222
+ )
223
+ new_probs = (
224
+ torch.topk(logits[:, -1:], k=1, dim=-1)
225
+ .values.squeeze(-1)
226
+ .to(torch.long)
227
+ )
228
+ batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
229
+ new_toks = torch.where(
230
+ batches_completed.unsqueeze(-1),
231
+ self.tokenizer.pad_token_id,
232
+ new_toks,
233
+ )
234
+ new_probs = torch.where(
235
+ batches_completed.unsqueeze(-1),
236
+ 0.0,
237
+ new_probs,
238
+ )
239
+ input_ids = torch.cat([input_ids, new_toks], dim=1)
240
+ input_probs = torch.cat([input_probs, new_probs], dim=1)
241
+
242
+ top_probs = input_probs.reshape(input_probs.size(0) // beam_size, beam_size, -1)
243
+ top_prob_inds = top_probs.sum(dim=-1).topk(k=1, dim=1).indices
244
+ winning_ids = torch.take_along_dim(
245
+ input_ids, top_prob_inds.unsqueeze(-1), 1
246
+ ).squeeze(1)
247
+ return self.tokenizer.decode(winning_ids, skip_special_tokens=True)
@@ -1,161 +0,0 @@
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)