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.
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/pyproject.toml +1 -1
- kernel_elastic_autoencoder-3.1.0/src/kernel_elastic_autoencoder/pipeline.py +247 -0
- kernel_elastic_autoencoder-3.0.0/src/kernel_elastic_autoencoder/pipeline.py +0 -161
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/README.md +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/config.py +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/losses.py +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/model.py +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -0
- {kernel_elastic_autoencoder-3.0.0 → kernel_elastic_autoencoder-3.1.0}/src/kernel_elastic_autoencoder/training.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernel_elastic_autoencoder"
|
|
3
|
-
version = "3.
|
|
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)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|