OneScience-Group/DiffDock
053
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 