Team Ai
Modelpublic

OneScience-Group/CodonTransformer

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes7downloads
pretrain.py284 linesDownload Raw Back to scripts
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