kernel-elastic-autoencoder 3.1.2__tar.gz → 3.2.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.1.2 → kernel_elastic_autoencoder-3.2.0}/PKG-INFO +1 -1
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/pyproject.toml +2 -1
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/config.py +3 -3
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/layers.py +2 -2
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/losses.py +1 -1
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/model.py +6 -9
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/pipeline.py +109 -64
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/training.py +6 -4
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/README.md +0 -0
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/__init__.py +0 -0
- {kernel_elastic_autoencoder-3.1.2 → kernel_elastic_autoencoder-3.2.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernel_elastic_autoencoder"
|
|
3
|
-
version = "3.
|
|
3
|
+
version = "3.2.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" }
|
|
@@ -45,4 +45,5 @@ dev = [
|
|
|
45
45
|
"pytest-xdist (>=3.8.0,<4.0.0)",
|
|
46
46
|
"python-semantic-release (>=10.6.1,<11.0.0)",
|
|
47
47
|
"pdoc (>=16.0.0,<17.0.0)",
|
|
48
|
+
"gitpython (>=3.1.0, <3.1.60)",
|
|
48
49
|
]
|
|
@@ -9,8 +9,8 @@ from pydantic import (
|
|
|
9
9
|
Field,
|
|
10
10
|
FilePath,
|
|
11
11
|
ImportString,
|
|
12
|
+
NonNegativeFloat,
|
|
12
13
|
NonNegativeInt,
|
|
13
|
-
PositiveFloat,
|
|
14
14
|
PositiveInt,
|
|
15
15
|
)
|
|
16
16
|
|
|
@@ -110,7 +110,7 @@ class ModelEncoderConfig(Config):
|
|
|
110
110
|
description="Factor by which the dimension of the hidden layer in the FFNs differs from the dimension of the "
|
|
111
111
|
"input in the encoder. Applies to Transformer and Compression FFNs.",
|
|
112
112
|
)
|
|
113
|
-
dropout:
|
|
113
|
+
dropout: NonNegativeFloat = Field(
|
|
114
114
|
default=0.0,
|
|
115
115
|
ge=0.0,
|
|
116
116
|
le=1.0,
|
|
@@ -136,7 +136,7 @@ class ModelDecoderConfig(Config):
|
|
|
136
136
|
description="Factor by which the dimension of the hidden layer in the FFNs differs from the dimension of the "
|
|
137
137
|
"input in the decoder. Applies to Transformer and Mixing FFNs.",
|
|
138
138
|
)
|
|
139
|
-
dropout:
|
|
139
|
+
dropout: NonNegativeFloat = Field(
|
|
140
140
|
default=0.1,
|
|
141
141
|
ge=0.0,
|
|
142
142
|
le=1.0,
|
|
@@ -31,7 +31,7 @@ class ConditionEmbedding(nn.Module):
|
|
|
31
31
|
|
|
32
32
|
def forward(self, c: torch.Tensor):
|
|
33
33
|
indices = self.indices.repeat(c.size(0), 1).to(c.device) # type: ignore
|
|
34
|
-
masked_indices = indices.masked_fill(c == self.padding_value, 0)
|
|
34
|
+
masked_indices = indices.masked_fill(c == self.padding_value, 0).to(c.device)
|
|
35
35
|
embedding = self.embedding(masked_indices)
|
|
36
36
|
return embedding * c.unsqueeze(-1).repeat(1, 1, embedding.size(-1))
|
|
37
37
|
|
|
@@ -205,7 +205,7 @@ class Decoder(nn.Module):
|
|
|
205
205
|
self.out_linear = nn.Linear(embedding_dim, vocab_size)
|
|
206
206
|
|
|
207
207
|
self.register_buffer(
|
|
208
|
-
"causal_mask", nn.Transformer.generate_square_subsequent_mask(max_len
|
|
208
|
+
"causal_mask", nn.Transformer.generate_square_subsequent_mask(max_len).to(torch.bool)
|
|
209
209
|
)
|
|
210
210
|
|
|
211
211
|
def forward(
|
|
@@ -95,7 +95,7 @@ class Loss(nn.Module):
|
|
|
95
95
|
prediction_noise: torch.Tensor,
|
|
96
96
|
ground_truth: torch.Tensor,
|
|
97
97
|
) -> torch.Tensor:
|
|
98
|
-
ground_truth.to(torch.long)
|
|
98
|
+
ground_truth = ground_truth.to(torch.long)
|
|
99
99
|
log_softmax = torch.log_softmax(prediction, dim=-1)
|
|
100
100
|
log_softmax_noise = torch.log_softmax(prediction_noise, dim=-1)
|
|
101
101
|
Y = F.one_hot(ground_truth, num_classes=prediction.size(-1))
|
|
@@ -108,13 +108,13 @@ class Model(
|
|
|
108
108
|
input_ids[:, :-1],
|
|
109
109
|
latents,
|
|
110
110
|
condition_embeddings,
|
|
111
|
-
padding_mask[: self.config_typed.input.max_len - 1],
|
|
111
|
+
padding_mask[:, : self.config_typed.input.max_len - 1],
|
|
112
112
|
)
|
|
113
113
|
prediction_noise = self.decoder(
|
|
114
114
|
input_ids[:, :-1],
|
|
115
115
|
latents_noise,
|
|
116
116
|
condition_embeddings,
|
|
117
|
-
padding_mask[: self.config_typed.input.max_len - 1],
|
|
117
|
+
padding_mask[:, : self.config_typed.input.max_len - 1],
|
|
118
118
|
)
|
|
119
119
|
return prediction, prediction_noise, latents_noise
|
|
120
120
|
|
|
@@ -139,8 +139,8 @@ class Model(
|
|
|
139
139
|
padding_mask = (
|
|
140
140
|
torch.cat((token_mask, condition_mask), 1)
|
|
141
141
|
if (token_mask is not None)
|
|
142
|
-
else torch.cat((torch.full_like(input_ids,
|
|
143
|
-
)
|
|
142
|
+
else torch.cat((torch.full_like(input_ids, False), condition_mask), 1)
|
|
143
|
+
).to(torch.bool)
|
|
144
144
|
return self.encoder(input_ids, conditions, padding_mask)[0]
|
|
145
145
|
|
|
146
146
|
def decode(
|
|
@@ -149,7 +149,6 @@ class Model(
|
|
|
149
149
|
latents: torch.Tensor,
|
|
150
150
|
condition_embeddings: torch.Tensor,
|
|
151
151
|
token_mask: torch.Tensor | None,
|
|
152
|
-
condition_mask: torch.Tensor,
|
|
153
152
|
) -> torch.Tensor:
|
|
154
153
|
"""Basic interface for a forward pass through the Model.model.decoder module.
|
|
155
154
|
|
|
@@ -161,15 +160,13 @@ class Model(
|
|
|
161
160
|
be produced using Model.embed_conditions.
|
|
162
161
|
token_mask: Tensor of dimension (B, S) containing boolean padding masks for each sequence. If None is
|
|
163
162
|
passed, the absence of padding is assumed.
|
|
164
|
-
condition_mask: Tensor of dimension (B, C) containing boolean condition padding masks for each sequence.
|
|
165
163
|
|
|
166
164
|
Returns:
|
|
167
165
|
torch.Tensor: Tensor of dimension (B, S, L) containing prediction logits produced by the decoder.
|
|
168
166
|
"""
|
|
169
|
-
padding_mask = (
|
|
170
|
-
torch.cat((token_mask, condition_mask), 1)
|
|
167
|
+
padding_mask = (token_mask
|
|
171
168
|
if (token_mask is not None)
|
|
172
|
-
else torch.
|
|
169
|
+
else torch.full_like(current_output, False)
|
|
173
170
|
).to(torch.bool)
|
|
174
171
|
return self.decoder(current_output, latents, condition_embeddings, padding_mask)
|
|
175
172
|
|
|
@@ -37,20 +37,21 @@ class Pipeline:
|
|
|
37
37
|
from HuggingFace Hub.
|
|
38
38
|
device: Torch device used for inference.
|
|
39
39
|
"""
|
|
40
|
-
self.
|
|
40
|
+
self.device = device
|
|
41
|
+
"""Torch device used for inference."""
|
|
42
|
+
if self.device is None:
|
|
43
|
+
self.device = torch.device("cpu")
|
|
44
|
+
self.model = model.to(self.device)
|
|
41
45
|
"""Pre-trained Model object for inference. Moved to Pipeline.device, and placed in eval() mode."""
|
|
42
46
|
self.model.eval()
|
|
43
47
|
self.tokenizer = tokenizer
|
|
44
48
|
"""Pre-configured Tokenizer object."""
|
|
45
|
-
self.device = device
|
|
46
|
-
"""Torch device used for inference."""
|
|
47
49
|
|
|
48
50
|
def completion(
|
|
49
51
|
self,
|
|
50
52
|
latents: torch.Tensor,
|
|
51
53
|
sequences: list[str],
|
|
52
54
|
conditions: list[list[float]] | torch.Tensor,
|
|
53
|
-
device: torch.device | None = None,
|
|
54
55
|
**kwargs,
|
|
55
56
|
) -> Iterable[str]:
|
|
56
57
|
"""Completes each conditioned input sequence.
|
|
@@ -62,7 +63,6 @@ class Pipeline:
|
|
|
62
63
|
latents: Tensor of dimension (B, P * E) containing latent vectors for the batch.
|
|
63
64
|
sequences: List of text sequences to complete.
|
|
64
65
|
conditions: List of condition value lists per batch.
|
|
65
|
-
device: Torch device used for inference.
|
|
66
66
|
**kwargs: Additional keyword arguments passed to Tokenizer.encode.
|
|
67
67
|
|
|
68
68
|
Returns:
|
|
@@ -70,38 +70,38 @@ class Pipeline:
|
|
|
70
70
|
"""
|
|
71
71
|
input_ids = self.tokenizer.encode(
|
|
72
72
|
text=sequences,
|
|
73
|
-
padding=
|
|
74
|
-
max_length=self.model.config_typed.input.max_len,
|
|
73
|
+
padding="longest",
|
|
75
74
|
add_special_tokens=False,
|
|
76
75
|
return_tensors="pt",
|
|
77
76
|
**kwargs,
|
|
78
77
|
).to(self.device)
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
|
|
78
|
+
input_ids = torch.cat(
|
|
79
|
+
[
|
|
80
|
+
torch.full(
|
|
81
|
+
(input_ids.size(0), 1),
|
|
82
|
+
self.tokenizer.bos_token_id,
|
|
83
|
+
device=self.device,
|
|
84
|
+
),
|
|
85
|
+
input_ids,
|
|
86
|
+
],
|
|
87
|
+
dim=1,
|
|
88
|
+
)
|
|
89
|
+
latents = latents.to(self.device)
|
|
90
|
+
conditions = torch.as_tensor(conditions, dtype=torch.float, device=self.device)
|
|
83
91
|
conds_embed = self.model.embed_conditions(conditions)
|
|
84
|
-
|
|
92
|
+
token_mask = (input_ids == self.tokenizer.pad_token_id).to(torch.bool)
|
|
93
|
+
batches_completed = torch.zeros(
|
|
94
|
+
input_ids.size(0), dtype=torch.bool, device=self.device
|
|
95
|
+
)
|
|
85
96
|
|
|
86
97
|
while (input_ids.size(1) < self.model.config_typed.input.max_len) and (
|
|
87
98
|
not batches_completed.all()
|
|
88
99
|
):
|
|
89
|
-
input_ids = torch.cat(
|
|
90
|
-
[
|
|
91
|
-
torch.full(
|
|
92
|
-
(input_ids.size(0), 1),
|
|
93
|
-
self.tokenizer.bos_token_id,
|
|
94
|
-
),
|
|
95
|
-
input_ids,
|
|
96
|
-
],
|
|
97
|
-
dim=1,
|
|
98
|
-
)
|
|
99
100
|
logits = self.model.decode(
|
|
100
101
|
current_output=input_ids,
|
|
101
102
|
latents=latents,
|
|
102
103
|
condition_embeddings=conds_embed,
|
|
103
|
-
token_mask=
|
|
104
|
-
condition_mask=condition_mask,
|
|
104
|
+
token_mask=token_mask,
|
|
105
105
|
)
|
|
106
106
|
new_toks = (
|
|
107
107
|
torch.topk(logits[:, -1:], k=1, dim=-1)
|
|
@@ -115,6 +115,15 @@ class Pipeline:
|
|
|
115
115
|
new_toks,
|
|
116
116
|
)
|
|
117
117
|
input_ids = torch.cat([input_ids, new_toks], dim=1)
|
|
118
|
+
token_mask = torch.cat(
|
|
119
|
+
[
|
|
120
|
+
token_mask,
|
|
121
|
+
torch.zeros(
|
|
122
|
+
input_ids.size(0), 1, dtype=torch.bool, device=self.device
|
|
123
|
+
),
|
|
124
|
+
],
|
|
125
|
+
dim=1,
|
|
126
|
+
)
|
|
118
127
|
return self.tokenizer.decode(input_ids, skip_special_tokens=True)
|
|
119
128
|
|
|
120
129
|
def beam_completion(
|
|
@@ -123,7 +132,6 @@ class Pipeline:
|
|
|
123
132
|
beam_size: int,
|
|
124
133
|
sequences: list[str],
|
|
125
134
|
conditions: list[list[float]] | torch.Tensor,
|
|
126
|
-
device: torch.device | None = None,
|
|
127
135
|
**kwargs,
|
|
128
136
|
) -> Iterable[str]:
|
|
129
137
|
"""Completes each conditioned input sequence, using the beam search strategy.
|
|
@@ -138,7 +146,6 @@ class Pipeline:
|
|
|
138
146
|
beam_size: Beam size of first step.
|
|
139
147
|
sequences: List of text sequences to complete.
|
|
140
148
|
conditions: List of condition value lists per batch.
|
|
141
|
-
device: Torch device used for inference.
|
|
142
149
|
**kwargs: Additional keyword arguments passed to Tokenizer.encode.
|
|
143
150
|
|
|
144
151
|
Returns:
|
|
@@ -146,42 +153,48 @@ class Pipeline:
|
|
|
146
153
|
"""
|
|
147
154
|
input_ids = self.tokenizer.encode(
|
|
148
155
|
text=sequences,
|
|
149
|
-
padding=
|
|
150
|
-
max_length=self.model.config_typed.input.max_len,
|
|
156
|
+
padding="longest",
|
|
151
157
|
add_special_tokens=False,
|
|
152
158
|
return_tensors="pt",
|
|
153
159
|
**kwargs,
|
|
154
160
|
).to(self.device)
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
|
|
158
|
-
|
|
159
|
-
|
|
160
|
-
|
|
161
|
+
input_ids = torch.cat(
|
|
162
|
+
[
|
|
163
|
+
torch.full(
|
|
164
|
+
(input_ids.size(0), 1),
|
|
165
|
+
self.tokenizer.bos_token_id,
|
|
166
|
+
device=self.device,
|
|
167
|
+
),
|
|
168
|
+
input_ids,
|
|
169
|
+
],
|
|
170
|
+
dim=1,
|
|
161
171
|
)
|
|
162
|
-
|
|
163
|
-
|
|
172
|
+
token_mask = (input_ids == self.tokenizer.pad_token_id).to(torch.bool)
|
|
173
|
+
latents = latents.to(self.device)
|
|
174
|
+
input_probs = torch.zeros_like(input_ids).repeat_interleave(beam_size, dim=0)
|
|
175
|
+
conditions = torch.as_tensor(conditions, dtype=torch.float, device=self.device)
|
|
176
|
+
conds_embed = self.model.embed_conditions(conditions)
|
|
177
|
+
batches_completed = torch.zeros(
|
|
178
|
+
input_ids.size(0) * beam_size, dtype=torch.bool, device=self.device
|
|
164
179
|
)
|
|
165
|
-
batches_completed = torch.zeros(input_ids.size(0) * beam_size, dtype=torch.bool)
|
|
166
180
|
|
|
167
181
|
logits = self.model.decode(
|
|
168
182
|
current_output=input_ids,
|
|
169
183
|
latents=latents,
|
|
170
184
|
condition_embeddings=conds_embed,
|
|
171
|
-
token_mask=
|
|
172
|
-
condition_mask=condition_mask,
|
|
185
|
+
token_mask=token_mask,
|
|
173
186
|
)
|
|
174
187
|
new_toks = (
|
|
175
188
|
torch.topk(logits[:, -1:], k=beam_size, dim=-1)
|
|
176
|
-
.indices.
|
|
189
|
+
.indices.flatten()
|
|
190
|
+
.unsqueeze(-1)
|
|
177
191
|
.to(torch.long)
|
|
178
192
|
)
|
|
179
193
|
new_probs = (
|
|
180
194
|
torch.topk(logits[:, -1:], k=beam_size, dim=-1)
|
|
181
|
-
.values.
|
|
182
|
-
.
|
|
195
|
+
.values.flatten()
|
|
196
|
+
.unsqueeze(-1)
|
|
183
197
|
)
|
|
184
|
-
batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
185
198
|
new_toks = torch.where(
|
|
186
199
|
batches_completed.unsqueeze(-1),
|
|
187
200
|
self.tokenizer.pad_token_id,
|
|
@@ -196,38 +209,36 @@ class Pipeline:
|
|
|
196
209
|
[input_ids.repeat_interleave(beam_size, dim=0), new_toks], dim=1
|
|
197
210
|
)
|
|
198
211
|
input_probs = torch.cat([input_probs, new_probs], dim=1)
|
|
212
|
+
conds_embed = conds_embed.repeat_interleave(beam_size, dim=0)
|
|
213
|
+
latents = latents.repeat_interleave(beam_size, dim=0)
|
|
214
|
+
token_mask = token_mask.repeat_interleave(beam_size, dim=0)
|
|
215
|
+
batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
216
|
+
token_mask = torch.cat(
|
|
217
|
+
[
|
|
218
|
+
token_mask,
|
|
219
|
+
torch.zeros(input_ids.size(0), 1, dtype=torch.bool, device=self.device),
|
|
220
|
+
],
|
|
221
|
+
dim=1,
|
|
222
|
+
)
|
|
199
223
|
|
|
200
224
|
while (input_ids.size(1) < self.model.config_typed.input.max_len) and (
|
|
201
225
|
not batches_completed.all()
|
|
202
226
|
):
|
|
203
|
-
input_ids = torch.cat(
|
|
204
|
-
[
|
|
205
|
-
torch.full(
|
|
206
|
-
(input_ids.size(0), 1),
|
|
207
|
-
self.tokenizer.bos_token_id,
|
|
208
|
-
),
|
|
209
|
-
input_ids,
|
|
210
|
-
],
|
|
211
|
-
dim=1,
|
|
212
|
-
)
|
|
213
227
|
logits = self.model.decode(
|
|
214
228
|
current_output=input_ids,
|
|
215
229
|
latents=latents,
|
|
216
230
|
condition_embeddings=conds_embed,
|
|
217
|
-
token_mask=
|
|
218
|
-
condition_mask=condition_mask,
|
|
231
|
+
token_mask=token_mask,
|
|
219
232
|
)
|
|
220
233
|
new_toks = (
|
|
221
234
|
torch.topk(logits[:, -1:], k=1, dim=-1)
|
|
222
|
-
.indices.
|
|
235
|
+
.indices.flatten()
|
|
236
|
+
.unsqueeze(-1)
|
|
223
237
|
.to(torch.long)
|
|
224
238
|
)
|
|
225
239
|
new_probs = (
|
|
226
|
-
torch.topk(logits[:, -1:], k=1, dim=-1)
|
|
227
|
-
.values.squeeze(-1)
|
|
228
|
-
.to(torch.long)
|
|
240
|
+
torch.topk(logits[:, -1:], k=1, dim=-1).values.flatten().unsqueeze(-1)
|
|
229
241
|
)
|
|
230
|
-
batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
231
242
|
new_toks = torch.where(
|
|
232
243
|
batches_completed.unsqueeze(-1),
|
|
233
244
|
self.tokenizer.pad_token_id,
|
|
@@ -240,10 +251,44 @@ class Pipeline:
|
|
|
240
251
|
)
|
|
241
252
|
input_ids = torch.cat([input_ids, new_toks], dim=1)
|
|
242
253
|
input_probs = torch.cat([input_probs, new_probs], dim=1)
|
|
254
|
+
token_mask = torch.cat(
|
|
255
|
+
[
|
|
256
|
+
token_mask,
|
|
257
|
+
torch.zeros(
|
|
258
|
+
input_ids.size(0), 1, dtype=torch.bool, device=self.device
|
|
259
|
+
),
|
|
260
|
+
],
|
|
261
|
+
dim=1,
|
|
262
|
+
)
|
|
263
|
+
batches_completed |= new_toks.squeeze(-1) == self.tokenizer.eos_token_id
|
|
243
264
|
|
|
244
265
|
top_probs = input_probs.reshape(input_probs.size(0) // beam_size, beam_size, -1)
|
|
245
|
-
top_prob_inds = top_probs.sum(dim=-1).topk(k=1, dim=1).indices
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
).squeeze(1)
|
|
266
|
+
top_prob_inds = top_probs.sum(dim=-1).topk(k=1, dim=1).indices.squeeze(-1)
|
|
267
|
+
grouped_ids = input_ids.view(top_probs.shape[0], beam_size, -1)
|
|
268
|
+
winning_ids = grouped_ids[torch.arange(top_probs.size(0)), top_prob_inds]
|
|
249
269
|
return self.tokenizer.decode(winning_ids, skip_special_tokens=True)
|
|
270
|
+
|
|
271
|
+
def encoding(
|
|
272
|
+
self, sequences: list[str], conditions: list[list[float]] | torch.Tensor
|
|
273
|
+
) -> torch.Tensor:
|
|
274
|
+
input_ids = self.tokenizer.encode(
|
|
275
|
+
text=sequences,
|
|
276
|
+
padding="max_length",
|
|
277
|
+
max_length=self.model.config_typed.input.max_len,
|
|
278
|
+
add_special_tokens=True,
|
|
279
|
+
return_tensors="pt",
|
|
280
|
+
).to(self.device)
|
|
281
|
+
conditions = torch.as_tensor(conditions, dtype=torch.float, device=self.device)
|
|
282
|
+
token_mask = (input_ids == self.model.config_typed.common.padding_idx).to(
|
|
283
|
+
dtype=torch.bool, device=self.device
|
|
284
|
+
)
|
|
285
|
+
condition_mask = (
|
|
286
|
+
conditions == self.model.config_typed.common.padding_value
|
|
287
|
+
).to(dtype=torch.bool, device=self.device)
|
|
288
|
+
|
|
289
|
+
return self.model.encode(
|
|
290
|
+
input_ids=input_ids,
|
|
291
|
+
conditions=conditions,
|
|
292
|
+
token_mask=token_mask,
|
|
293
|
+
condition_mask=condition_mask,
|
|
294
|
+
)
|
|
@@ -76,8 +76,8 @@ class Trainer:
|
|
|
76
76
|
return_tensors="pt",
|
|
77
77
|
)
|
|
78
78
|
conditions = torch.as_tensor(conditions, dtype=torch.float)
|
|
79
|
-
token_mask = (input_ids
|
|
80
|
-
condition_mask = (conditions
|
|
79
|
+
token_mask = (input_ids == model.config_typed.common.padding_idx).to(torch.bool)
|
|
80
|
+
condition_mask = (conditions == model.config_typed.common.padding_value).to(
|
|
81
81
|
torch.bool
|
|
82
82
|
)
|
|
83
83
|
|
|
@@ -127,7 +127,8 @@ class Trainer:
|
|
|
127
127
|
)
|
|
128
128
|
accelerator.backward(loss)
|
|
129
129
|
optimizer.step()
|
|
130
|
-
|
|
130
|
+
print(f"Train loss: {float(loss.detach())}")
|
|
131
|
+
|
|
131
132
|
model.eval()
|
|
132
133
|
for input_ids, conditions, token_mask, condition_mask in tqdm(
|
|
133
134
|
dataloader_test, desc=f"Epoch {epoch}, Test Batch"
|
|
@@ -142,8 +143,9 @@ class Trainer:
|
|
|
142
143
|
loss = loss_fn(
|
|
143
144
|
prediction, prediction_noise, input_ids[:, 1:], latents_noise
|
|
144
145
|
)
|
|
146
|
+
print(f"Test loss: {float(loss.detach())}")
|
|
145
147
|
|
|
146
|
-
scheduler.step(
|
|
148
|
+
scheduler.step()
|
|
147
149
|
|
|
148
150
|
accelerator.wait_for_everyone()
|
|
149
151
|
curr_epoch += 1
|
|
File without changes
|
|
File without changes
|
|
File without changes
|