""" File: pretrain.py ------------------- Pretrain the CodonTransformer model. The dataset is a JSON file. You can use prepare_training_data from CodonData to prepare the dataset. The repository README has a guide on how to prepare the dataset and use this script. """ import argparse import gzip import math import os import sys from pathlib import Path PROJECT_ROOT = Path(__file__).resolve().parents[1] MODEL_DIR = PROJECT_ROOT / "model" if str(MODEL_DIR) not in sys.path: sys.path.insert(0, str(MODEL_DIR)) import pytorch_lightning as pl import torch from torch.utils.data import DataLoader from transformers import BigBirdConfig, BigBirdForMaskedLM, PreTrainedTokenizerFast from CodonTransformer.CodonUtils import ( MAX_LEN, NUM_ORGANISMS, TOKEN2MASK, IterableJSONData, ) class MaskedTokenizerCollator: def __init__(self, tokenizer): self.tokenizer = tokenizer def __call__(self, examples): tokenized = self.tokenizer( [ex["codons"] for ex in examples], return_attention_mask=True, return_token_type_ids=True, truncation=True, padding=True, max_length=MAX_LEN, return_tensors="pt", ) seq_len = tokenized["input_ids"].shape[-1] species_index = torch.tensor([[ex["organism"]] for ex in examples]) tokenized["token_type_ids"] = species_index.repeat(1, seq_len) inputs = tokenized["input_ids"] targets = inputs.clone() prob_matrix = torch.full(inputs.shape, 0.15) prob_matrix[inputs < 5] = 0.0 selected = torch.bernoulli(prob_matrix).bool() # 80% of the time, replace masked input tokens with respective mask tokens replaced = torch.bernoulli(torch.full(selected.shape, 0.8)).bool() & selected inputs[replaced] = torch.tensor( list((map(TOKEN2MASK.__getitem__, inputs[replaced].numpy()))) ) # 10% of the time, we replace masked input tokens with random vector. randomized = ( torch.bernoulli(torch.full(selected.shape, 0.1)).bool() & selected & ~replaced ) random_idx = torch.randint(26, 90, inputs.shape, dtype=torch.long) inputs[randomized] = random_idx[randomized] tokenized["input_ids"] = inputs tokenized["labels"] = torch.where(selected, targets, -100) return tokenized class plTrainHarness(pl.LightningModule): def __init__(self, model, learning_rate, warmup_fraction, total_training_steps): super().__init__() self.model = model self.learning_rate = learning_rate self.warmup_fraction = warmup_fraction self.total_training_steps = total_training_steps def configure_optimizers(self): optimizer = torch.optim.AdamW( self.model.parameters(), lr=self.learning_rate, ) total_steps = self.total_training_steps or self.trainer.estimated_stepping_batches if total_steps <= 0: raise ValueError(f"Expected positive integer total_steps, but got {total_steps}") lr_scheduler = { "scheduler": torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr=self.learning_rate, total_steps=total_steps, pct_start=self.warmup_fraction, ), "interval": "step", "frequency": 1, } return [optimizer], [lr_scheduler] def training_step(self, batch, batch_idx): self.model.bert.set_attention_type("block_sparse") outputs = self.model(**batch) self.log_dict( dictionary={ "loss": outputs.loss, "lr": self.trainer.optimizers[0].param_groups[0]["lr"], }, on_step=True, prog_bar=True, ) return outputs.loss class EpochCheckpoint(pl.Callback): def __init__(self, checkpoint_dir, save_interval): super().__init__() self.checkpoint_dir = checkpoint_dir self.save_interval = save_interval def on_train_epoch_end(self, trainer, pl_module): current_epoch = trainer.current_epoch if current_epoch % self.save_interval == 0 or current_epoch == 0: checkpoint_path = os.path.join( self.checkpoint_dir, f"epoch_{current_epoch}.ckpt" ) trainer.save_checkpoint(checkpoint_path) print(f"\nCheckpoint saved at {checkpoint_path}\n") def count_jsonl_records(path): open_fn = gzip.open if path.endswith(".gz") else open with open_fn(path, "rt") as file: return sum(1 for line in file if line.strip()) def estimate_training_steps(args): num_records = count_jsonl_records(args.train_data_path) num_devices = 1 if args.debug else args.num_gpus samples_per_step = max(1, args.batch_size * num_devices) batches_per_epoch = math.ceil(num_records / samples_per_step) optimizer_steps_per_epoch = math.ceil( batches_per_epoch / max(1, args.accumulate_grad_batches) ) total_steps = max(1, optimizer_steps_per_epoch * args.max_epochs) print( "Estimated training steps: " f"{total_steps} " f"({num_records} records, batch_size={args.batch_size}, " f"devices={num_devices}, max_epochs={args.max_epochs}, " f"accumulate_grad_batches={args.accumulate_grad_batches})" ) return total_steps def main(args): """Pretrain the CodonTransformer model.""" pl.seed_everything(args.seed) torch.set_float32_matmul_precision("medium") total_training_steps = estimate_training_steps(args) # Load the tokenizer and model tokenizer = PreTrainedTokenizerFast( tokenizer_file=args.tokenizer_path, bos_token="[CLS]", eos_token="[SEP]", unk_token="[UNK]", sep_token="[SEP]", pad_token="[PAD]", cls_token="[CLS]", mask_token="[MASK]", ) config = BigBirdConfig( vocab_size=len(tokenizer), type_vocab_size=NUM_ORGANISMS, sep_token_id=2, ) model = BigBirdForMaskedLM(config=config) harnessed_model = plTrainHarness( model, args.learning_rate, args.warmup_fraction, total_training_steps, ) # Load the training data train_data = IterableJSONData(args.train_data_path, dist_env="slurm") data_loader = DataLoader( dataset=train_data, collate_fn=MaskedTokenizerCollator(tokenizer), batch_size=args.batch_size, num_workers=0 if args.debug else args.num_workers, persistent_workers=False if args.debug else True, ) # Setup trainer and callbacks save_checkpoint = EpochCheckpoint(args.checkpoint_dir, args.save_interval) trainer = pl.Trainer( default_root_dir=args.checkpoint_dir, strategy="ddp_find_unused_parameters_true", accelerator="gpu", devices=1 if args.debug else args.num_gpus, precision="16-mixed", max_epochs=args.max_epochs, deterministic=False, enable_checkpointing=True, callbacks=[save_checkpoint], accumulate_grad_batches=args.accumulate_grad_batches, ) # Pretrain the model trainer.fit(harnessed_model, data_loader) if __name__ == "__main__": parser = argparse.ArgumentParser(description="Pretrain the CodonTransformer model.") parser.add_argument( "--tokenizer_path", type=str, required=True, help="Path to the tokenizer model file", ) parser.add_argument( "--train_data_path", type=str, required=True, help="Path to the training data JSON file", ) parser.add_argument( "--checkpoint_dir", type=str, required=True, help="Directory where checkpoints will be saved", ) parser.add_argument( "--batch_size", type=int, default=6, help="Batch size for training" ) parser.add_argument( "--max_epochs", type=int, default=5, help="Maximum number of epochs to train" ) parser.add_argument( "--num_workers", type=int, default=5, help="Number of workers for data loading" ) parser.add_argument( "--accumulate_grad_batches", type=int, default=1, help="Number of batches to accumulate gradients", ) parser.add_argument( "--num_gpus", type=int, default=16, help="Number of GPUs to use for training" ) parser.add_argument( "--learning_rate", type=float, default=5e-5, help="Learning rate for the optimizer", ) parser.add_argument( "--warmup_fraction", type=float, default=0.1, help="Fraction of total steps to use for warmup", ) parser.add_argument( "--save_interval", type=int, default=5, help="Save checkpoint every N epochs" ) parser.add_argument( "--seed", type=int, default=123, help="Random seed for reproducibility" ) parser.add_argument("--debug", action="store_true", help="Enable debug mode") args = parser.parse_args() main(args)