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()