kernel-elastic-autoencoder 1.0.0__tar.gz → 3.0.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.
Files changed (15) hide show
  1. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/PKG-INFO +2 -2
  2. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/pyproject.toml +5 -3
  3. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/__init__.py +11 -20
  4. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/losses.py +1 -1
  5. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/model.py +9 -46
  6. kernel_elastic_autoencoder-3.0.0/src/kernel_elastic_autoencoder/pipeline.py +161 -0
  7. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/tokenizer.py +0 -68
  8. kernel_elastic_autoencoder-3.0.0/src/kernel_elastic_autoencoder/training.py +153 -0
  9. kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/collate.py +0 -164
  10. kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/pipeline.py +0 -203
  11. kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/sample.py +0 -58
  12. kernel_elastic_autoencoder-1.0.0/src/kernel_elastic_autoencoder/training.py +0 -386
  13. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/README.md +0 -0
  14. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/config.py +0 -0
  15. {kernel_elastic_autoencoder-1.0.0 → kernel_elastic_autoencoder-3.0.0}/src/kernel_elastic_autoencoder/layers.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernel_elastic_autoencoder
3
- Version: 1.0.0
3
+ Version: 3.0.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
@@ -11,8 +11,8 @@ Classifier: Programming Language :: Python :: 3
11
11
  Classifier: Programming Language :: Python :: 3.12
12
12
  Classifier: Programming Language :: Python :: 3.13
13
13
  Classifier: Programming Language :: Python :: 3.14
14
+ Requires-Dist: accelerate (>=1.14.0,<2.0.0)
14
15
  Requires-Dist: huggingface-hub (>=1.22.0,<2.0.0)
15
- Requires-Dist: pandas (>=3.0.3,<4.0.0)
16
16
  Requires-Dist: pydantic (>=2.13.4,<3.0.0)
17
17
  Requires-Dist: safetensors (>=0.8.0,<0.9.0)
18
18
  Description-Content-Type: text/markdown
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "kernel_elastic_autoencoder"
3
- version = "1.0.0"
3
+ version = "3.0.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" }
@@ -11,13 +11,16 @@ requires-python = ">=3.12,<3.15"
11
11
  dependencies = [
12
12
  "pydantic (>=2.13.4,<3.0.0)",
13
13
  "huggingface-hub (>=1.22.0,<2.0.0)",
14
- "pandas (>=3.0.3,<4.0.0)",
15
14
  "safetensors (>=0.8.0,<0.9.0)",
15
+ "accelerate (>=1.14.0,<2.0.0)",
16
16
  ]
17
17
 
18
18
  [tool.poetry]
19
19
  packages = [{ include = "kernel_elastic_autoencoder", from = "src" }]
20
20
 
21
+ [tool.ruff.lint]
22
+ ignore = ["B008"]
23
+
21
24
  [tool.semantic_release]
22
25
  version_toml = ["pyproject.toml:project.version"]
23
26
  commit_parser = "conventional"
@@ -40,7 +43,6 @@ dev = [
40
43
  "ruff (>=0.15.20,<0.17.0)",
41
44
  "pytest-cov (>=7.1.0,<8.0.0)",
42
45
  "pytest-xdist (>=3.8.0,<4.0.0)",
43
- "ty (>=0.0.62,<0.0.64)",
44
46
  "python-semantic-release (>=10.6.1,<11.0.0)",
45
47
  "pdoc (>=16.0.0,<17.0.0)",
46
48
  ]
@@ -1,4 +1,3 @@
1
- from kernel_elastic_autoencoder.collate import Collated, Collator, DataframeCollator
2
1
  from kernel_elastic_autoencoder.config import (
3
2
  ExperimentConfig,
4
3
  ModelCommonConfig,
@@ -13,32 +12,24 @@ from kernel_elastic_autoencoder.config import (
13
12
  )
14
13
  from kernel_elastic_autoencoder.losses import Loss
15
14
  from kernel_elastic_autoencoder.model import Model
16
- from kernel_elastic_autoencoder.pipeline import Completion, Pipeline
17
- from kernel_elastic_autoencoder.sample import Sampler, Top1Sampler
15
+ from kernel_elastic_autoencoder.pipeline import Pipeline
18
16
  from kernel_elastic_autoencoder.tokenizer import Tokenizer
19
- from kernel_elastic_autoencoder.training import Trainer, TrainerCallback
17
+ from kernel_elastic_autoencoder.training import Trainer
20
18
 
21
19
  __all__ = [
22
20
  "ExperimentConfig",
23
- "ModelConfig",
21
+ "Loss",
22
+ "Model",
24
23
  "ModelCommonConfig",
25
- "ModelInputConfig",
26
- "ModelEncoderConfig",
24
+ "ModelConfig",
27
25
  "ModelDecoderConfig",
28
- "TrainingConfig",
29
- "TrainingCommonConfig",
30
- "TrainingHyperparameterConfig",
31
- "TrainingOptimizerConfig",
26
+ "ModelEncoderConfig",
27
+ "ModelInputConfig",
32
28
  "Pipeline",
33
- "Completion",
34
- "Model",
35
- "Loss",
36
29
  "Tokenizer",
37
- "Collator",
38
- "Collated",
39
- "DataframeCollator",
40
- "Sampler",
41
- "Top1Sampler",
42
30
  "Trainer",
43
- "TrainerCallback",
31
+ "TrainingCommonConfig",
32
+ "TrainingConfig",
33
+ "TrainingHyperparameterConfig",
34
+ "TrainingOptimizerConfig",
44
35
  ]
@@ -121,7 +121,7 @@ class Loss(nn.Module):
121
121
  / (2 * (self.hp_sigma**2))
122
122
  ).sum()
123
123
  loss = self.hp_lambda * (
124
- torch.tensor([1])
124
+ torch.tensor(1, dtype=latents.dtype, device=latents.device)
125
125
  - ((1 / (latents.size(0) * self.kernel_dist_size)) * kernel_pairwise_sum)
126
126
  )
127
127
  return loss
@@ -1,9 +1,5 @@
1
- from pathlib import Path
2
- from typing import Any
3
-
4
1
  import torch
5
2
  from huggingface_hub import PyTorchModelHubMixin
6
- from huggingface_hub.hub_mixin import T, DataclassInstance
7
3
  from torch import nn
8
4
 
9
5
  from kernel_elastic_autoencoder.config import ModelConfig
@@ -78,47 +74,6 @@ class Model(
78
74
  )
79
75
  """Model decoder module."""
80
76
 
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
77
  def forward(
123
78
  self,
124
79
  input_ids: torch.Tensor,
@@ -218,5 +173,13 @@ class Model(
218
173
  ).to(torch.bool)
219
174
  return self.decoder(current_output, latents, condition_embeddings, padding_mask)
220
175
 
221
- def embed_conditions(self, conditions):
176
+ def embed_conditions(self, conditions: torch.Tensor) -> torch.Tensor:
177
+ """Basic interface for the separate embedding of condition vectors.
178
+
179
+ Args:
180
+ conditions: Tensor of dimension (B, C) containing condition values for each sequence.
181
+
182
+ Returns:
183
+ torch.Tensor: Tensor of dimension (B, C, E) containing condition embeddings for each sequence.
184
+ """
222
185
  return self.encoder.embedding.conditional_embedding(conditions)
@@ -0,0 +1,161 @@
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)
@@ -78,71 +78,3 @@ class Tokenizer(Protocol):
78
78
  Tokenizer: Pretrained tokenizer.
79
79
  """
80
80
  ...
81
-
82
-
83
- class _DummySlowTokenizer(Tokenizer):
84
- def __init__(self, train: Iterable[str]):
85
- self.pad_token = "_"
86
- self.bos_token = "?"
87
- self.eos_token = "!"
88
-
89
- vocab = set[str]()
90
- for b in train:
91
- vocab = vocab.union(set(b))
92
-
93
- special_tokens = [self.pad_token, self.bos_token, self.eos_token]
94
- special_tokens.extend(list(vocab))
95
- self.vocab = special_tokens
96
- self.vocab_size = len(self.vocab)
97
-
98
- self.pad_token_id = 0
99
- self.bos_token_id = 1
100
- self.eos_token_id = 2
101
-
102
- def encode(
103
- self,
104
- seq: Iterable[str],
105
- padding: bool,
106
- max_length: int,
107
- add_special_tokens: bool,
108
- **kwargs,
109
- ) -> torch.Tensor:
110
- ids_col = list[torch.Tensor]()
111
- for b in seq:
112
- toks = list(b)
113
- if add_special_tokens:
114
- toks.insert(0, self.bos_token)
115
- toks.append(self.eos_token)
116
- if padding:
117
- toks += [self.pad_token] * (max_length - len(toks))
118
- ids = torch.tensor([self.vocab.index(t) for t in toks], dtype=torch.long)
119
- ids_col.append(ids)
120
- return torch.stack(ids_col, dim=0)
121
-
122
- def decode(
123
- self, ids: torch.Tensor, skip_special_tokens: bool, **kwargs
124
- ) -> Iterable[str]:
125
- if skip_special_tokens:
126
- return [
127
- str(
128
- [
129
- (
130
- self.vocab[idx]
131
- if idx
132
- not in [
133
- self.pad_token_id,
134
- self.bos_token_id,
135
- self.eos_token_id,
136
- ]
137
- else ""
138
- )
139
- for idx in b
140
- ]
141
- )
142
- for b in ids
143
- ]
144
- else:
145
- return [str([self.vocab[idx] for idx in b]) for b in ids]
146
-
147
- @classmethod
148
- def from_pretrained(cls, pretrained_model_name_or_path: str | Path, **kwargs): ...
@@ -0,0 +1,153 @@
1
+ import os
2
+ from collections.abc import Iterable
3
+
4
+ import torch
5
+ from accelerate import Accelerator
6
+ from accelerate.utils import tqdm
7
+
8
+ from kernel_elastic_autoencoder.config import TrainingConfig
9
+ from kernel_elastic_autoencoder.losses import Loss
10
+ from kernel_elastic_autoencoder.model import Model
11
+ from kernel_elastic_autoencoder.tokenizer import Tokenizer
12
+
13
+
14
+ class Trainer:
15
+ def __init__(
16
+ self,
17
+ config: dict | TrainingConfig,
18
+ ):
19
+ """Instantiates a Trainer object.
20
+
21
+ Args:
22
+ config: Dictionary or TrainingConfig schema defining model parameters. Will be validated with TrainingConfig
23
+ regardless of input type.
24
+ """
25
+ self.config = config
26
+ """Configuration object for Hugging Face Hub compatible serialization. Not recommended to use, as it is
27
+ internal-use. Use Trainer.config_typed instead."""
28
+ config_typed = TrainingConfig.model_validate(config)
29
+ self.config_typed: TrainingConfig = config_typed
30
+ """Type-validated config in a Pydantic TrainingConfig schema, recommended for public API use."""
31
+
32
+ def train(
33
+ self,
34
+ model: Model,
35
+ tokenizer: Tokenizer,
36
+ sequences: Iterable[str],
37
+ conditions: Iterable[Iterable[float]] | torch.Tensor,
38
+ train_split: float = 0.9,
39
+ checkpoint: str = "./checkpoint",
40
+ ):
41
+ """Trains a model. Optionally, resumes from an existing checkpoint.
42
+
43
+ Args:
44
+ model: Freshly instantiated model.
45
+ tokenizer: Pre-configured Tokenizer object. The Tokenizer protocol supports tokenizers
46
+ inheriting from transformers.PreTrainedTokenizerBase, so such tokenizers may be loaded
47
+ from HuggingFace Hub.
48
+ sequences: Iterable of text sequences.
49
+ conditions: Iterable of iterables of condition values per sequence.
50
+ train_split: Fraction of dataset used for training. Must be between 0 and 1.
51
+ checkpoint: Path of local checkpoint to be saved and/or resumed.
52
+ """
53
+ accelerator = Accelerator()
54
+
55
+ loss_fn = Loss(
56
+ hp_lambda=self.config_typed.hyperparameters.hp_lambda,
57
+ hp_delta=self.config_typed.hyperparameters.hp_delta,
58
+ hp_sigma=self.config_typed.hyperparameters.hp_sigma,
59
+ kernel_dist_size=self.config_typed.hyperparameters.kernel_dist_size,
60
+ padding_idx=model.config_typed.common.padding_idx,
61
+ embedding_dim=model.config_typed.common.embedding_dim,
62
+ pooling_dim=model.config_typed.common.pooling_dim,
63
+ )
64
+ optimizer = self.config_typed.optimizer.optimizer_fn(
65
+ model.parameters(), **self.config_typed.optimizer.optimizer_params
66
+ )
67
+ scheduler = self.config_typed.optimizer.scheduler_fn(
68
+ optimizer, **self.config_typed.optimizer.scheduler_params
69
+ )
70
+
71
+ input_ids = tokenizer.encode(
72
+ seq=sequences,
73
+ padding=True,
74
+ max_length=model.config_typed.input.max_len,
75
+ add_special_tokens=True,
76
+ )
77
+ conditions = torch.as_tensor(conditions, dtype=torch.float)
78
+ token_mask = (input_ids != model.config_typed.common.padding_idx).to(torch.bool)
79
+ condition_mask = (conditions != model.config_typed.common.padding_value).to(
80
+ torch.bool
81
+ )
82
+
83
+ dataset = torch.utils.data.TensorDataset(
84
+ input_ids, conditions, token_mask, condition_mask
85
+ )
86
+ dataset_train, dataset_test = torch.utils.data.random_split(
87
+ dataset, [train_split, 1 - train_split]
88
+ )
89
+ dataloader_train = torch.utils.data.DataLoader(
90
+ dataset_train,
91
+ batch_size=self.config_typed.common.batch_size,
92
+ )
93
+ dataloader_test = torch.utils.data.DataLoader(
94
+ dataset_test,
95
+ batch_size=self.config_typed.common.batch_size,
96
+ )
97
+ curr_epoch = 0
98
+
99
+ model, optimizer, dataloader_train, dataloader_test, scheduler, curr_epoch = (
100
+ accelerator.prepare(
101
+ model,
102
+ optimizer,
103
+ dataloader_train,
104
+ dataloader_test,
105
+ scheduler,
106
+ curr_epoch,
107
+ )
108
+ )
109
+ accelerator.register_for_checkpointing(scheduler, curr_epoch)
110
+
111
+ if os.path.exists(checkpoint):
112
+ accelerator.load_state(checkpoint)
113
+
114
+ for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
115
+ model.train()
116
+ for input_ids, conditions, token_mask, condition_mask in tqdm(
117
+ dataloader_train, desc=f"Epoch {epoch}, Train Batch"
118
+ ):
119
+ optimizer.zero_grad()
120
+ prediction, prediction_noise, latents_noise = model(
121
+ input_ids, conditions, token_mask, condition_mask
122
+ )
123
+ loss = loss_fn(
124
+ prediction, prediction_noise, input_ids[:, 1:], latents_noise
125
+ )
126
+ accelerator.backward(loss)
127
+ optimizer.step()
128
+
129
+ model.eval()
130
+ for input_ids, conditions, token_mask, condition_mask in tqdm(
131
+ dataloader_test, desc=f"Epoch {epoch}, Test Batch"
132
+ ):
133
+ with torch.no_grad():
134
+ prediction, prediction_noise, latents_noise = model(
135
+ input_ids,
136
+ conditions,
137
+ token_mask,
138
+ condition_mask,
139
+ )
140
+ loss = loss_fn(
141
+ prediction, prediction_noise, input_ids[:, 1:], latents_noise
142
+ )
143
+
144
+ scheduler.step(epoch)
145
+
146
+ accelerator.wait_for_everyone()
147
+ curr_epoch += 1
148
+ os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
149
+ model.save_pretrained(os.path.join(checkpoint, "dist/"))
150
+ self.config_typed.to_json(
151
+ os.path.join(checkpoint, "dist/train_config.json")
152
+ )
153
+ accelerator.save_state(checkpoint)