OneScience-Group/CodonTransformer
07
1"""2File: pretrain.py3-------------------4Pretrain the CodonTransformer model.5 6The dataset is a JSON file. You can use prepare_training_data from CodonData to7prepare the dataset. The repository README has a guide on how to prepare the8dataset and use this script.9"""10 11import argparse12import gzip13import math14import os15import sys16from pathlib import Path17 18PROJECT_ROOT = Path(__file__).resolve().parents[1]19MODEL_DIR = PROJECT_ROOT / "model"20if str(MODEL_DIR) not in sys.path:21 sys.path.insert(0, str(MODEL_DIR))22 23import pytorch_lightning as pl24import torch25from torch.utils.data import DataLoader26from transformers import BigBirdConfig, BigBirdForMaskedLM, PreTrainedTokenizerFast27 28from CodonTransformer.CodonUtils import (29 MAX_LEN,30 NUM_ORGANISMS,31 TOKEN2MASK,32 IterableJSONData,33)34 35 36class MaskedTokenizerCollator:37 def __init__(self, tokenizer):38 self.tokenizer = tokenizer39 40 def __call__(self, examples):41 tokenized = self.tokenizer(42 [ex["codons"] for ex in examples],43 return_attention_mask=True,44 return_token_type_ids=True,45 truncation=True,46 padding=True,47 max_length=MAX_LEN,48 return_tensors="pt",49 )50 51 seq_len = tokenized["input_ids"].shape[-1]52 species_index = torch.tensor([[ex["organism"]] for ex in examples])53 tokenized["token_type_ids"] = species_index.repeat(1, seq_len)54 55 inputs = tokenized["input_ids"]56 targets = inputs.clone()57 58 prob_matrix = torch.full(inputs.shape, 0.15)59 prob_matrix[inputs < 5] = 0.060 selected = torch.bernoulli(prob_matrix).bool()61 62 # 80% of the time, replace masked input tokens with respective mask tokens63 replaced = torch.bernoulli(torch.full(selected.shape, 0.8)).bool() & selected64 inputs[replaced] = torch.tensor(65 list((map(TOKEN2MASK.__getitem__, inputs[replaced].numpy())))66 )67 68 # 10% of the time, we replace masked input tokens with random vector.69 randomized = (70 torch.bernoulli(torch.full(selected.shape, 0.1)).bool()71 & selected72 & ~replaced73 )74 random_idx = torch.randint(26, 90, inputs.shape, dtype=torch.long)75 inputs[randomized] = random_idx[randomized]76 77 tokenized["input_ids"] = inputs78 tokenized["labels"] = torch.where(selected, targets, -100)79 80 return tokenized81 82 83class plTrainHarness(pl.LightningModule):84 def __init__(self, model, learning_rate, warmup_fraction, total_training_steps):85 super().__init__()86 self.model = model87 self.learning_rate = learning_rate88 self.warmup_fraction = warmup_fraction89 self.total_training_steps = total_training_steps90 91 def configure_optimizers(self):92 optimizer = torch.optim.AdamW(93 self.model.parameters(),94 lr=self.learning_rate,95 )96 total_steps = self.total_training_steps or self.trainer.estimated_stepping_batches97 if total_steps <= 0:98 raise ValueError(f"Expected positive integer total_steps, but got {total_steps}")99 lr_scheduler = {100 "scheduler": torch.optim.lr_scheduler.OneCycleLR(101 optimizer,102 max_lr=self.learning_rate,103 total_steps=total_steps,104 pct_start=self.warmup_fraction,105 ),106 "interval": "step",107 "frequency": 1,108 }109 return [optimizer], [lr_scheduler]110 111 def training_step(self, batch, batch_idx):112 self.model.bert.set_attention_type("block_sparse")113 outputs = self.model(**batch)114 self.log_dict(115 dictionary={116 "loss": outputs.loss,117 "lr": self.trainer.optimizers[0].param_groups[0]["lr"],118 },119 on_step=True,120 prog_bar=True,121 )122 return outputs.loss123 124 125class EpochCheckpoint(pl.Callback):126 def __init__(self, checkpoint_dir, save_interval):127 super().__init__()128 self.checkpoint_dir = checkpoint_dir129 self.save_interval = save_interval130 131 def on_train_epoch_end(self, trainer, pl_module):132 current_epoch = trainer.current_epoch133 if current_epoch % self.save_interval == 0 or current_epoch == 0:134 checkpoint_path = os.path.join(135 self.checkpoint_dir, f"epoch_{current_epoch}.ckpt"136 )137 trainer.save_checkpoint(checkpoint_path)138 print(f"\nCheckpoint saved at {checkpoint_path}\n")139 140 141def count_jsonl_records(path):142 open_fn = gzip.open if path.endswith(".gz") else open143 with open_fn(path, "rt") as file:144 return sum(1 for line in file if line.strip())145 146 147def estimate_training_steps(args):148 num_records = count_jsonl_records(args.train_data_path)149 num_devices = 1 if args.debug else args.num_gpus150 samples_per_step = max(1, args.batch_size * num_devices)151 batches_per_epoch = math.ceil(num_records / samples_per_step)152 optimizer_steps_per_epoch = math.ceil(153 batches_per_epoch / max(1, args.accumulate_grad_batches)154 )155 total_steps = max(1, optimizer_steps_per_epoch * args.max_epochs)156 print(157 "Estimated training steps: "158 f"{total_steps} "159 f"({num_records} records, batch_size={args.batch_size}, "160 f"devices={num_devices}, max_epochs={args.max_epochs}, "161 f"accumulate_grad_batches={args.accumulate_grad_batches})"162 )163 return total_steps164 165 166def main(args):167 """Pretrain the CodonTransformer model."""168 pl.seed_everything(args.seed)169 torch.set_float32_matmul_precision("medium")170 total_training_steps = estimate_training_steps(args)171 172 # Load the tokenizer and model173 tokenizer = PreTrainedTokenizerFast(174 tokenizer_file=args.tokenizer_path,175 bos_token="[CLS]",176 eos_token="[SEP]",177 unk_token="[UNK]",178 sep_token="[SEP]",179 pad_token="[PAD]",180 cls_token="[CLS]",181 mask_token="[MASK]",182 )183 config = BigBirdConfig(184 vocab_size=len(tokenizer),185 type_vocab_size=NUM_ORGANISMS,186 sep_token_id=2,187 )188 model = BigBirdForMaskedLM(config=config)189 harnessed_model = plTrainHarness(190 model,191 args.learning_rate,192 args.warmup_fraction,193 total_training_steps,194 )195 196 # Load the training data197 train_data = IterableJSONData(args.train_data_path, dist_env="slurm")198 data_loader = DataLoader(199 dataset=train_data,200 collate_fn=MaskedTokenizerCollator(tokenizer),201 batch_size=args.batch_size,202 num_workers=0 if args.debug else args.num_workers,203 persistent_workers=False if args.debug else True,204 )205 206 # Setup trainer and callbacks207 save_checkpoint = EpochCheckpoint(args.checkpoint_dir, args.save_interval)208 trainer = pl.Trainer(209 default_root_dir=args.checkpoint_dir,210 strategy="ddp_find_unused_parameters_true",211 accelerator="gpu",212 devices=1 if args.debug else args.num_gpus,213 precision="16-mixed",214 max_epochs=args.max_epochs,215 deterministic=False,216 enable_checkpointing=True,217 callbacks=[save_checkpoint],218 accumulate_grad_batches=args.accumulate_grad_batches,219 )220 221 # Pretrain the model222 trainer.fit(harnessed_model, data_loader)223 224 225if __name__ == "__main__":226 parser = argparse.ArgumentParser(description="Pretrain the CodonTransformer model.")227 parser.add_argument(228 "--tokenizer_path",229 type=str,230 required=True,231 help="Path to the tokenizer model file",232 )233 parser.add_argument(234 "--train_data_path",235 type=str,236 required=True,237 help="Path to the training data JSON file",238 )239 parser.add_argument(240 "--checkpoint_dir",241 type=str,242 required=True,243 help="Directory where checkpoints will be saved",244 )245 parser.add_argument(246 "--batch_size", type=int, default=6, help="Batch size for training"247 )248 parser.add_argument(249 "--max_epochs", type=int, default=5, help="Maximum number of epochs to train"250 )251 parser.add_argument(252 "--num_workers", type=int, default=5, help="Number of workers for data loading"253 )254 parser.add_argument(255 "--accumulate_grad_batches",256 type=int,257 default=1,258 help="Number of batches to accumulate gradients",259 )260 parser.add_argument(261 "--num_gpus", type=int, default=16, help="Number of GPUs to use for training"262 )263 parser.add_argument(264 "--learning_rate",265 type=float,266 default=5e-5,267 help="Learning rate for the optimizer",268 )269 parser.add_argument(270 "--warmup_fraction",271 type=float,272 default=0.1,273 help="Fraction of total steps to use for warmup",274 )275 parser.add_argument(276 "--save_interval", type=int, default=5, help="Save checkpoint every N epochs"277 )278 parser.add_argument(279 "--seed", type=int, default=123, help="Random seed for reproducibility"280 )281 parser.add_argument("--debug", action="store_true", help="Enable debug mode")282 args = parser.parse_args()283 main(args)284 