Team Ai
Modelpublic

OneScience-Group/DiffDock

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes53downloads
train_diffdock.py405 linesDownload Raw Back to scripts
1import sys2from pathlib import Path3 4DIR = Path(__file__).resolve().parent.parent5sys.path.insert(0, str(DIR))6 7import argparse8import copy9import math10import os11import random12from functools import partial13from pathlib import Path14from types import SimpleNamespace15 16import numpy as np17import torch18import yaml19 20from onescience.datapipes.diffdock.loader import construct_loader21from onescience.utils.diffdock.diffusion_utils import t_to_sigma as t_to_sigma_compl22from onescience.utils.diffdock.training import (23    inference_epoch_fix,24    loss_function,25    test_epoch,26    train_epoch,27)28from onescience.utils.diffdock.utils import (29    ExponentialMovingAverage,30    get_optimizer_and_scheduler,31    save_yaml_file,32)33from onescience.utils.diffdock.validation import validate_training_entrypoint34 35try:36    from models.score_wrapper import build_score_model37except ImportError:38    from models.score_wrapper import build_score_model39 40 41def parse_args():42    parser = argparse.ArgumentParser()43    parser.add_argument("--config", required=True, help="Path to the training YAML config.")44    return parser.parse_args()45 46def _resolve_env_vars(obj):47    if isinstance(obj, str):48        return os.path.expandvars(obj)49    if isinstance(obj, dict):50        return {k: _resolve_env_vars(v) for k, v in obj.items()}51    if isinstance(obj, list):52        return [_resolve_env_vars(v) for v in obj]53    return obj54 55# def load_config(config_path):56#     with open(config_path, "r", encoding="utf-8") as handle:57#         return yaml.safe_load(handle) or {}58def load_config(config_path):59    with open(config_path, "r", encoding="utf-8") as handle:60        return _resolve_env_vars(yaml.safe_load(handle) or {})61 62 63def flatten_config(config):64    flat = {}65    for key, value in config.items():66        if isinstance(value, dict):67            flat.update(value)68        else:69            flat[key] = value70    return flat71 72 73def to_namespace(config):74    return SimpleNamespace(**config)75 76 77def resolve_device(device_name):78    if device_name in {None, "auto"}:79        return torch.device("cuda:0" if torch.cuda.is_available() else "cpu")80    return torch.device(device_name)81 82 83def state_dict_for_save(model):84    return model.module.state_dict() if hasattr(model, "module") else model.state_dict()85 86 87def model_target(model):88    return model.module if hasattr(model, "module") else model89 90 91def set_seed(seed):92    random.seed(seed)93    np.random.seed(seed)94    torch.manual_seed(seed)95    if torch.cuda.is_available():96        torch.cuda.manual_seed_all(seed)97 98 99def maybe_init_wandb(args):100    if not args.wandb:101        return None102    try:103        import wandb104    except ImportError as exc:105        raise ImportError("wandb is enabled in the config but is not installed.") from exc106    wandb.init(project=args.project, name=args.run_name, config=vars(args))107    return wandb108 109 110def maybe_load_restart(args, model, optimizer, ema_weights):111    if args.restart_dir is None:112        return113 114    checkpoint_path = Path(args.restart_dir) / f"{args.restart_ckpt}.pt"115    try:116        checkpoint = torch.load(checkpoint_path, map_location=torch.device("cpu"))117        if args.restart_lr is not None:118            checkpoint["optimizer"]["param_groups"][0]["lr"] = args.restart_lr119        optimizer.load_state_dict(checkpoint["optimizer"])120        model_target(model).load_state_dict(checkpoint["model"], strict=True)121        ema_weights.load_state_dict(checkpoint["ema_weights"], device=args.device)122        print("Restarting from epoch", checkpoint["epoch"])123    except Exception as exc:124        print("Exception", exc)125        checkpoint = torch.load(Path(args.restart_dir) / "best_model.pt", map_location=torch.device("cpu"))126        model_target(model).load_state_dict(checkpoint, strict=True)127        print("Due to exception had to take the best epoch and no optimiser")128 129 130def maybe_load_pretrain(args, model):131    if args.pretrain_dir is None:132        return133    checkpoint = torch.load(134        Path(args.pretrain_dir) / f"{args.pretrain_ckpt}.pt",135        map_location=torch.device("cpu"),136    )137    if isinstance(checkpoint, dict) and "model" in checkpoint and "optimizer" in checkpoint:138        checkpoint = checkpoint["model"]139    model_target(model).load_state_dict(checkpoint, strict=True)140    print("Using pretrained model", str(Path(args.pretrain_dir) / f"{args.pretrain_ckpt}.pt"))141 142 143def train(args, model, optimizer, scheduler, ema_weights, train_loader, val_loader, t_to_sigma, run_dir, val_dataset2):144    loss_fn = partial(145        loss_function,146        tr_weight=args.tr_weight,147        rot_weight=args.rot_weight,148        tor_weight=args.tor_weight,149        no_torsion=args.no_torsion,150        backbone_weight=args.backbone_loss_weight,151        sidechain_weight=args.sidechain_loss_weight,152    )153 154    best_val_loss = math.inf155    best_val_inference_value = math.inf if args.inference_earlystop_goal == "min" else 0156    best_val_secondary_value = math.inf if args.inference_earlystop_goal == "min" else 0157    best_epoch = 0158    best_val_inference_epoch = 0159 160    freeze_params = 0161    scheduler_mode = args.inference_earlystop_goal if args.val_inference_freq is not None else "min"162    if args.scheduler == "layer_linear_warmup":163        freeze_params = args.warmup_dur * (args.num_conv_layers + 2) - 1164        print("Freezing some parameters until epoch {}".format(freeze_params))165 166    wandb_run = maybe_init_wandb(args)167 168    print("Starting training...")169    for epoch in range(args.n_epochs):170        if epoch % 5 == 0:171            print("Run name:", args.run_name)172 173        if args.scheduler == "layer_linear_warmup" and (epoch + 1) % args.warmup_dur == 0:174            step = (epoch + 1) // args.warmup_dur175            if step < args.num_conv_layers + 2:176                print("New unfreezing step")177                optimizer, scheduler = get_optimizer_and_scheduler(178                    args,179                    model,180                    step=step,181                    scheduler_mode=scheduler_mode,182                )183            elif step == args.num_conv_layers + 2:184                print("Unfreezing all parameters")185                optimizer, scheduler = get_optimizer_and_scheduler(186                    args,187                    model,188                    step=step,189                    scheduler_mode=scheduler_mode,190                )191                ema_weights = ExponentialMovingAverage(model.parameters(), decay=args.ema_rate)192        elif args.scheduler == "linear_warmup" and epoch == args.warmup_dur:193            print("Moving to plateu scheduler")194            optimizer, scheduler = get_optimizer_and_scheduler(195                args,196                model,197                step=1,198                scheduler_mode=scheduler_mode,199                optimizer=optimizer,200            )201 202        logs = {}203        train_losses = train_epoch(204            model,205            train_loader,206            optimizer,207            args.device,208            t_to_sigma,209            loss_fn,210            ema_weights if epoch > freeze_params else None,211        )212        print(213            "Epoch {}: Training loss {:.4f}  tr {:.4f}   rot {:.4f}   tor {:.4f}   sc {:.4f}  lr {:.4f}".format(214                epoch,215                train_losses["loss"],216                train_losses["tr_loss"],217                train_losses["rot_loss"],218                train_losses["tor_loss"],219                train_losses["sidechain_loss"],220                optimizer.param_groups[0]["lr"],221            )222        )223 224        if epoch > freeze_params:225            ema_weights.store(model.parameters())226            if args.use_ema:227                ema_weights.copy_to(model.parameters())228 229        val_losses = test_epoch(model, val_loader, args.device, t_to_sigma, loss_fn, args.test_sigma_intervals)230        print(231            "Epoch {}: Validation loss {:.4f}  tr {:.4f}   rot {:.4f}   tor {:.4f}   sc {:.4f}".format(232                epoch,233                val_losses["loss"],234                val_losses["tr_loss"],235                val_losses["rot_loss"],236                val_losses["tor_loss"],237                val_losses["sidechain_loss"],238            )239        )240 241        if args.val_inference_freq is not None and (epoch + 1) % args.val_inference_freq == 0:242            inf_dataset = [243                val_loader.dataset.get(i)244                for i in range(min(args.num_inference_complexes, val_loader.dataset.__len__()))245            ]246            inf_metrics = inference_epoch_fix(model, inf_dataset, args.device, t_to_sigma, args)247            print(248                "Epoch {}: Val inference rmsds_lt2 {:.3f} rmsds_lt5 {:.3f} min_rmsds_lt2 {:.3f} min_rmsds_lt5 {:.3f}".format(249                    epoch,250                    inf_metrics["rmsds_lt2"],251                    inf_metrics["rmsds_lt5"],252                    inf_metrics["min_rmsds_lt2"],253                    inf_metrics["min_rmsds_lt5"],254                )255            )256            logs.update({"valinf_" + k: v for k, v in inf_metrics.items()})257 258        if args.double_val and args.val_inference_freq is not None and (epoch + 1) % args.val_inference_freq == 0:259            inf_dataset = [260                val_dataset2.get(i)261                for i in range(min(args.num_inference_complexes, val_dataset2.__len__()))262            ]263            inf_metrics2 = inference_epoch_fix(model, inf_dataset, args.device, t_to_sigma, args)264            print(265                "Epoch {}: Val inference on second validation rmsds_lt2 {:.3f} rmsds_lt5 {:.3f} min_rmsds_lt2 {:.3f} min_rmsds_lt5 {:.3f}".format(266                    epoch,267                    inf_metrics2["rmsds_lt2"],268                    inf_metrics2["rmsds_lt5"],269                    inf_metrics2["min_rmsds_lt2"],270                    inf_metrics2["min_rmsds_lt5"],271                )272            )273            logs.update({"valinf2_" + k: v for k, v in inf_metrics2.items()})274            logs.update({"valinfcomb_" + k: (v + inf_metrics[k]) / 2 for k, v in inf_metrics2.items()})275 276        if args.train_inference_freq is not None and (epoch + 1) % args.train_inference_freq == 0:277            inf_dataset = [278                train_loader.dataset.get(i)279                for i in range(min(min(args.num_inference_complexes, 300), train_loader.dataset.__len__()))280            ]281            inf_metrics = inference_epoch_fix(model, inf_dataset, args.device, t_to_sigma, args)282            print(283                "Epoch {}: Train inference rmsds_lt2 {:.3f} rmsds_lt5 {:.3f} min_rmsds_lt2 {:.3f} min_rmsds_lt5 {:.3f}".format(284                    epoch,285                    inf_metrics["rmsds_lt2"],286                    inf_metrics["rmsds_lt5"],287                    inf_metrics["min_rmsds_lt2"],288                    inf_metrics["min_rmsds_lt5"],289                )290            )291            logs.update({"traininf_" + k: v for k, v in inf_metrics.items()})292 293        if epoch > freeze_params:294            if not args.use_ema:295                ema_weights.copy_to(model.parameters())296            ema_state_dict = copy.deepcopy(state_dict_for_save(model))297            ema_weights.restore(model.parameters())298        else:299            ema_state_dict = copy.deepcopy(state_dict_for_save(model))300 301        if wandb_run is not None:302            logs.update({"train_" + k: v for k, v in train_losses.items()})303            logs.update({"val_" + k: v for k, v in val_losses.items()})304            logs["current_lr"] = optimizer.param_groups[0]["lr"]305            wandb_run.log(logs, step=epoch + 1)306 307        model_state_dict = state_dict_for_save(model)308        if args.inference_earlystop_metric in logs and (309            (args.inference_earlystop_goal == "min" and logs[args.inference_earlystop_metric] <= best_val_inference_value)310            or (args.inference_earlystop_goal == "max" and logs[args.inference_earlystop_metric] >= best_val_inference_value)311        ):312            best_val_inference_value = logs[args.inference_earlystop_metric]313            best_val_inference_epoch = epoch314            torch.save(model_state_dict, os.path.join(run_dir, "best_inference_epoch_model.pt"))315            if epoch > freeze_params:316                torch.save(ema_state_dict, os.path.join(run_dir, "best_ema_inference_epoch_model.pt"))317 318        if args.inference_secondary_metric is not None and args.inference_secondary_metric in logs and (319            (args.inference_earlystop_goal == "min" and logs[args.inference_secondary_metric] <= best_val_secondary_value)320            or (args.inference_earlystop_goal == "max" and logs[args.inference_secondary_metric] >= best_val_secondary_value)321        ):322            best_val_secondary_value = logs[args.inference_secondary_metric]323            if epoch > freeze_params:324                torch.save(ema_state_dict, os.path.join(run_dir, "best_ema_secondary_epoch_model.pt"))325 326        if val_losses["loss"] <= best_val_loss:327            best_val_loss = val_losses["loss"]328            best_epoch = epoch329            torch.save(model_state_dict, os.path.join(run_dir, "best_model.pt"))330            if epoch > freeze_params:331                torch.save(ema_state_dict, os.path.join(run_dir, "best_ema_model.pt"))332 333        if args.save_model_freq is not None and (epoch + 1) % args.save_model_freq == 0:334            best_model_path = os.path.join(run_dir, "best_model.pt")335            if os.path.exists(best_model_path):336                torch.save(torch.load(best_model_path, map_location=torch.device("cpu")), os.path.join(run_dir, f"epoch{epoch + 1}_best_model.pt"))337 338        if scheduler:339            if epoch < freeze_params or (args.scheduler == "linear_warmup" and epoch < args.warmup_dur):340                scheduler.step()341            elif args.val_inference_freq is not None:342                scheduler.step(best_val_inference_value)343            else:344                scheduler.step(val_losses["loss"])345 346        torch.save(347            {348                "epoch": epoch,349                "model": model_state_dict,350                "optimizer": optimizer.state_dict(),351                "ema_weights": ema_weights.state_dict(),352            },353            os.path.join(run_dir, "last_model.pt"),354        )355 356    print("Best Validation Loss {} on Epoch {}".format(best_val_loss, best_epoch))357    print("Best inference metric {} on Epoch {}".format(best_val_inference_value, best_val_inference_epoch))358 359 360def main():361    parsed = parse_args()362    raw_config = load_config(parsed.config)363    flat_config = flatten_config(raw_config)364    args = to_namespace(flat_config)365 366    if getattr(args, "run_name", None) in {None, ""}:367        args.run_name = Path(parsed.config).stem368 369    args.device = resolve_device(getattr(args, "device", "auto"))370    if getattr(args, "cudnn_benchmark", False) and args.device.type == "cuda":371        torch.backends.cudnn.benchmark = True372    set_seed(getattr(args, "seed", 0))373    validate_training_entrypoint(args)374 375    assert args.inference_earlystop_goal in {"max", "min"}376    if args.val_inference_freq is not None and args.scheduler is not None:377        assert args.scheduler_patience > args.val_inference_freq378 379    run_dir = os.path.join(args.log_dir, args.run_name)380    saved_args = vars(args).copy()381    saved_args["device"] = str(args.device)382    save_yaml_file(os.path.join(run_dir, "model_parameters.yml"), saved_args)383 384    t_to_sigma = partial(t_to_sigma_compl, args=args)385    train_loader, val_loader, val_dataset2 = construct_loader(args, t_to_sigma, args.device)386    model, _ = build_score_model(args, args.device, no_parallel=False)387    optimizer, scheduler = get_optimizer_and_scheduler(388        args,389        model,390        scheduler_mode=args.inference_earlystop_goal if args.val_inference_freq is not None else "min",391    )392    ema_weights = ExponentialMovingAverage(model.parameters(), decay=args.ema_rate)393 394    maybe_load_restart(args, model, optimizer, ema_weights)395    maybe_load_pretrain(args, model)396 397    numel = sum(p.numel() for p in model.parameters())398    print("Model with", numel, "parameters")399 400    train(args, model, optimizer, scheduler, ema_weights, train_loader, val_loader, t_to_sigma, run_dir, val_dataset2)401 402 403if __name__ == "__main__":404    main()405