kernel-elastic-autoencoder 3.2.1__tar.gz → 3.2.3__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.2.1
3
+ Version: 3.2.3
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.2.1"
3
+ version = "3.2.3"
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" }
@@ -3,13 +3,12 @@ from collections.abc import Iterable
3
3
 
4
4
  import torch
5
5
  from accelerate import Accelerator
6
- from accelerate.utils import tqdm
6
+ from accelerate.utils import DistributedDataParallelKwargs, tqdm
7
7
 
8
8
  from kernel_elastic_autoencoder.config import TrainingConfig
9
9
  from kernel_elastic_autoencoder.losses import Loss
10
10
  from kernel_elastic_autoencoder.model import Model
11
11
  from kernel_elastic_autoencoder.tokenizer import Tokenizer
12
- from accelerate.utils import DistributedDataParallelKwargs
13
12
 
14
13
 
15
14
  class Trainer:
@@ -116,11 +115,14 @@ class Trainer:
116
115
  accelerator.load_state(checkpoint)
117
116
  curr_epoch = scheduler.scheduler.last_epoch + 1
118
117
 
118
+ accelerator.wait_for_everyone()
119
119
  for epoch in range(curr_epoch, self.config_typed.common.max_epochs):
120
120
  model.train()
121
- for input_ids, conditions, token_mask, condition_mask in tqdm(
122
- dataloader_train, desc=f"Epoch {epoch}, Train Batch"
123
- ):
121
+ if accelerator.is_local_main_process:
122
+ batch_bar = tqdm(
123
+ total=len(dataloader_train), desc=f"Epoch {epoch}, Train Batch"
124
+ )
125
+ for input_ids, conditions, token_mask, condition_mask in dataloader_train:
124
126
  optimizer.zero_grad()
125
127
  prediction, prediction_noise, latents_noise = model(
126
128
  input_ids, conditions, token_mask, condition_mask
@@ -130,12 +132,16 @@ class Trainer:
130
132
  )
131
133
  accelerator.backward(loss)
132
134
  optimizer.step()
133
- print(f"Train loss: {float(loss.detach())}")
135
+ if accelerator.is_local_main_process:
136
+ batch_bar.update(1)
137
+ accelerator.print(f"Train loss: {float(loss.detach())}")
134
138
 
135
139
  model.eval()
136
- for input_ids, conditions, token_mask, condition_mask in tqdm(
137
- dataloader_test, desc=f"Epoch {epoch}, Test Batch"
138
- ):
140
+ if accelerator.is_local_main_process:
141
+ batch_bar = tqdm(
142
+ total=len(dataloader_train), desc=f"Epoch {epoch}, Test Batch"
143
+ )
144
+ for input_ids, conditions, token_mask, condition_mask in dataloader_test:
139
145
  with torch.no_grad():
140
146
  prediction, prediction_noise, latents_noise = model(
141
147
  input_ids,
@@ -146,15 +152,20 @@ class Trainer:
146
152
  loss = loss_fn(
147
153
  prediction, prediction_noise, input_ids[:, 1:], latents_noise
148
154
  )
149
- print(f"Test loss: {float(loss.detach())}")
155
+ if accelerator.is_local_main_process:
156
+ batch_bar.update(1)
157
+ accelerator.print(f"Test loss: {float(loss.detach())}")
150
158
 
151
159
  scheduler.step()
152
160
 
153
- accelerator.wait_for_everyone()
154
161
  curr_epoch += 1
155
- os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
156
- model.save_pretrained(os.path.join(checkpoint, "dist/"))
157
- self.config_typed.to_json(
158
- os.path.join(checkpoint, "dist/train_config.json")
159
- )
162
+ accelerator.wait_for_everyone()
160
163
  accelerator.save_state(checkpoint)
164
+ if accelerator.is_main_process:
165
+ os.makedirs(os.path.join(checkpoint, "dist/"), exist_ok=True)
166
+ accelerator.unwrap_model(model).save_pretrained(
167
+ os.path.join(checkpoint, "dist/")
168
+ )
169
+ self.config_typed.to_json(
170
+ os.path.join(checkpoint, "dist/train_config.json")
171
+ )