pCoMole / train.py
AlienChen's picture
Upload 83 files
7f316fe verified
Raw History Blame Contribute Delete
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()