Download train.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 5.8 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/train.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/train.py
-
curl -L -o train.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/train.py
5.8 kB
| import argparse | |
| import os | |
| import pytorch_lightning as pl | |
| import torch | |
| import yaml | |
| from datasets import load_from_disk | |
| from easydict import EasyDict as edict | |
| from pytorch_lightning.callbacks import LearningRateMonitor, ModelCheckpoint | |
| from torch.utils.data import DataLoader | |
| from transformers import EsmTokenizer | |
| from logic import flow | |
| from model.reparam_models import EditFlow, ProteinEditFlowModel, SMILESEditFlowModel | |
| from paths import add_repo_to_sys_path, resolve_path, selfies_vocab_path, smiles_vocab_files | |
| from smiles_tokenizer.my_tokenizers import SMILES_SPE_Tokenizer | |
| from smiles_tokenizer.selfies_tokenizers import SelfiesTokenizer | |
| add_repo_to_sys_path() | |
| def load_config(config_path: str) -> edict: | |
| with open(resolve_path(config_path), "r") as f: | |
| return edict(yaml.safe_load(f)) | |
| def build_editflow(cfg, device=None): | |
| if cfg.task == "protein": | |
| tokenizer = EsmTokenizer.from_pretrained("facebook/esm2_t33_650M_UR50D") | |
| vocab_size = 24 | |
| source_distribution = flow.get_source_distribution( | |
| source_distribution=cfg.flow.source_distribution, | |
| vocab_size=vocab_size, | |
| special_token_ids=[0, 1, 2, 3], | |
| ) | |
| pad_id, bos_id, eos_id = 1, 0, 2 | |
| model = ProteinEditFlowModel(vocab_size=vocab_size, pad_id=pad_id, config=cfg.model) | |
| elif cfg.task == "smiles": | |
| vocab_path, splits_path = smiles_vocab_files() | |
| tokenizer = SMILES_SPE_Tokenizer(str(vocab_path), str(splits_path)) | |
| vocab_size = 586 | |
| source_distribution = flow.get_source_distribution( | |
| source_distribution=cfg.flow.source_distribution, | |
| vocab_size=vocab_size, | |
| special_token_ids=[0, 1, 2, 3, 4], | |
| ) | |
| pad_id, bos_id, eos_id = 0, 2, 3 | |
| model = SMILESEditFlowModel(vocab_size=vocab_size, pad_id=pad_id, config=cfg.model) | |
| elif cfg.task == "selfies": | |
| tokenizer = SelfiesTokenizer.load(str(selfies_vocab_path())) | |
| vocab_size = 44 | |
| source_distribution = flow.get_source_distribution( | |
| source_distribution=cfg.flow.source_distribution, | |
| vocab_size=vocab_size, | |
| special_token_ids=[0, 1, 2], | |
| ) | |
| pad_id, bos_id, eos_id = 0, 1, 2 | |
| model = SMILESEditFlowModel(vocab_size=vocab_size, pad_id=pad_id, config=cfg.model) | |
| else: | |
| raise NotImplementedError(f"Unsupported task: {cfg.task}") | |
| if device is not None: | |
| model = model.to(device) | |
| eps_id = getattr(cfg.flow, "eps_id", -1) | |
| path = flow.get_path(scheduler_type=cfg.flow.scheduler_type, exponent=cfg.flow.exponent, eps_id=eps_id) | |
| loss_fn = flow.get_loss_function(loss_function=cfg.flow.loss_function, path=path) | |
| editflow = EditFlow(model, loss_fn, path, source_distribution, pad_id, bos_id, eos_id, cfg) | |
| return editflow, tokenizer, source_distribution, pad_id, bos_id, eos_id, eps_id | |
| def build_dataloaders(cfg): | |
| train_path = resolve_path(cfg.data.train_path) | |
| val_path = resolve_path(cfg.data.val_path) | |
| if not train_path.exists() or not val_path.exists(): | |
| raise FileNotFoundError( | |
| "Training data was not found.\n" | |
| f" train: {train_path}\n" | |
| f" val: {val_path}\n" | |
| "Update data.train_path / data.val_path in the config. " | |
| "The shipped SELFIES peptidomimetic dataset lives at data/selfies/28k_mimetics." | |
| ) | |
| num_workers = int(getattr(getattr(cfg, "data", {}), "num_workers", 4) or 4) | |
| train_dataloader = DataLoader(load_from_disk(str(train_path)), batch_size=None, shuffle=True, num_workers=num_workers) | |
| val_dataloader = DataLoader(load_from_disk(str(val_path)), batch_size=None, shuffle=False, num_workers=num_workers) | |
| return train_dataloader, val_dataloader | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Train an Edit Flow model") | |
| parser.add_argument("--config", type=str, required=True, help="Path to YAML config") | |
| parser.add_argument("--wandb", action="store_true", help="Log to Weights & Biases") | |
| args = parser.parse_args() | |
| cfg = load_config(args.config) | |
| run_name = ( | |
| f"reparam_{cfg.task}_lr{cfg.optim.lr}_epoch{cfg.optim.n_epochs}" | |
| f"_scale{cfg.model.scale_size}_optimal{cfg.model.p_optimal}" | |
| ) | |
| workdir = resolve_path(getattr(cfg, "work_dir", "outputs")) / run_name | |
| os.makedirs(workdir, exist_ok=True) | |
| pl.seed_everything(cfg.training.seed, workers=True) | |
| editflow, _, _, _, _, _, _ = build_editflow(cfg) | |
| train_dataloader, val_dataloader = build_dataloaders(cfg) | |
| ckpt = ModelCheckpoint( | |
| dirpath=os.path.join(workdir, "checkpoint"), | |
| monitor="val_loss", | |
| mode="min", | |
| save_top_k=3, | |
| save_last=True, | |
| filename="epoch{epoch:04d}-val{val_loss:.2f}", | |
| auto_insert_metric_name=False, | |
| ) | |
| callbacks = [ckpt, LearningRateMonitor(logging_interval="step")] | |
| logger = False | |
| if args.wandb: | |
| from pytorch_lightning.loggers import WandbLogger | |
| logger = WandbLogger( | |
| project=getattr(getattr(cfg, "logging", {}), "project", "pCoMole"), | |
| name=run_name, | |
| entity=getattr(getattr(cfg, "logging", {}), "entity", None), | |
| ) | |
| trainer = pl.Trainer( | |
| default_root_dir=str(workdir), | |
| accelerator="gpu" if torch.cuda.is_available() else "cpu", | |
| devices=cfg.compute.ngpus, | |
| strategy="ddp" if cfg.compute.ngpus > 1 else "auto", | |
| precision="bf16-mixed", | |
| max_epochs=cfg.optim.n_epochs, | |
| log_every_n_steps=10, | |
| callbacks=callbacks, | |
| enable_checkpointing=True, | |
| gradient_clip_val=1.0, | |
| deterministic=False, | |
| logger=logger, | |
| ) | |
| trainer.fit(editflow, train_dataloader, val_dataloader) | |
| if __name__ == "__main__": | |
| main() | |