kernel-elastic-autoencoder 3.1.3__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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 3.1.3
3
+ Version: 3.2.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.1.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" }
@@ -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: PositiveFloat = Field(
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: PositiveFloat = Field(
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 - 1).to(torch.bool)
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, True), condition_mask), 1)
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.cat((torch.full_like(current_output, False), condition_mask), 1)
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.model = model.to(device)
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='do_not_pad',
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
- conditions = torch.as_tensor(conditions, dtype=torch.float, device=device)
80
- condition_mask = (
81
- conditions != self.model.config_typed.common.padding_value
82
- ).to(torch.bool)
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
- batches_completed = torch.zeros(input_ids.size(0), dtype=torch.bool)
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=None,
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='do_not_pad',
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
- input_probs = torch.zeros_like(input_ids * beam_size)
156
- conditions = torch.as_tensor(conditions, dtype=torch.float, device=device)
157
- condition_mask = (
158
- (conditions != self.model.config_typed.common.padding_value)
159
- .to(torch.bool)
160
- .repeat_interleave(beam_size, dim=0)
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
- conds_embed = self.model.embed_conditions(conditions).repeat_interleave(
163
- beam_size, dim=0
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=None,
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.squeeze(-1)
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.squeeze(-1)
182
- .to(torch.long)
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=None,
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.squeeze(-1)
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
- winning_ids = torch.take_along_dim(
247
- input_ids, top_prob_inds.unsqueeze(-1), 1
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 != model.config_typed.common.padding_idx).to(torch.bool)
80
- condition_mask = (conditions != model.config_typed.common.padding_value).to(
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(epoch)
148
+ scheduler.step()
147
149
 
148
150
  accelerator.wait_for_everyone()
149
151
  curr_epoch += 1