Download scripts/finetune.py from OneScience-Group/CodonTransformer: direct link, hf CLI and curl.
- Browser
- Download file 8.99 kB
-
https://huggingface.co/OneScience-Group/CodonTransformer/resolve/main/scripts/finetune.py
- Command line
-
hf download hf://OneScience-Group/CodonTransformer/scripts/finetune.py
-
curl -L -o finetune.py https://huggingface.co/OneScience-Group/CodonTransformer/resolve/main/scripts/finetune.py
8.99 kB
| """ | |
| File: finetune.py | |
| ------------------- | |
| Finetune the CodonTransformer model. | |
| The pretrained model is loaded directly from Hugging Face. | |
| 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 AutoTokenizer, BigBirdForMaskedLM | |
| from CodonTransformer.CodonUtils import ( | |
| MAX_LEN, | |
| 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 = tokenized["input_ids"].clone() | |
| prob_matrix = torch.full(inputs.shape, 0.15) | |
| prob_matrix[torch.where(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, prob_matrix.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 DumpStateDict(pl.callbacks.ModelCheckpoint): | |
| def __init__(self, checkpoint_dir, checkpoint_filename, every_n_train_steps): | |
| super().__init__( | |
| dirpath=checkpoint_dir, every_n_train_steps=every_n_train_steps | |
| ) | |
| self.checkpoint_filename = checkpoint_filename | |
| def on_save_checkpoint(self, trainer, pl_module, checkpoint): | |
| model = pl_module.model | |
| torch.save( | |
| model.state_dict(), os.path.join(self.dirpath, self.checkpoint_filename) | |
| ) | |
| 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.dataset_dir) | |
| 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): | |
| """Finetune 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 = AutoTokenizer.from_pretrained("adibvafa/CodonTransformer") | |
| model = BigBirdForMaskedLM.from_pretrained("adibvafa/CodonTransformer-base") | |
| harnessed_model = plTrainHarness( | |
| model, | |
| args.learning_rate, | |
| args.warmup_fraction, | |
| total_training_steps, | |
| ) | |
| # Load the training data | |
| train_data = IterableJSONData(args.dataset_dir, 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 = DumpStateDict( | |
| checkpoint_dir=args.checkpoint_dir, | |
| checkpoint_filename=args.checkpoint_filename, | |
| every_n_train_steps=args.save_every_n_steps, | |
| ) | |
| 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, | |
| ) | |
| # Finetune the model | |
| trainer.fit(harnessed_model, data_loader) | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser(description="Finetune the CodonTransformer model.") | |
| parser.add_argument( | |
| "--dataset_dir", | |
| type=str, | |
| required=True, | |
| help="Directory containing the dataset", | |
| ) | |
| parser.add_argument( | |
| "--checkpoint_dir", | |
| type=str, | |
| required=True, | |
| help="Directory where checkpoints will be saved", | |
| ) | |
| parser.add_argument( | |
| "--checkpoint_filename", | |
| type=str, | |
| default="finetune.ckpt", | |
| help="Filename for the saved checkpoint", | |
| ) | |
| parser.add_argument( | |
| "--batch_size", type=int, default=6, help="Batch size for training" | |
| ) | |
| parser.add_argument( | |
| "--max_epochs", type=int, default=15, 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=4, 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_every_n_steps", | |
| type=int, | |
| default=512, | |
| help="Save checkpoint every N steps", | |
| ) | |
| 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) | |