import argparse import numpy as np import torch as th import torch.multiprocessing from torch_geometric.loader import DataLoader from onescience.datapipes.genscore.data import PDBbindDataset from models.model.model import GatedGCN, GenScore, GraphTransformer from onescience.metrics.genscore.utils import ( EarlyStopping, run_a_train_epoch, run_an_eval_epoch, set_random_seed, ) torch.multiprocessing.set_sharing_strategy("file_system") def parse_args(): parser = argparse.ArgumentParser(description="Train GenScore.") parser.add_argument("--num_epochs", type=int, default=5000) parser.add_argument("--batch_size", type=int, default=64) parser.add_argument("--aux_weight", type=float, default=0.001) parser.add_argument("--affi_weight", type=float, default=-0.5) parser.add_argument("--patience", type=int, default=70) parser.add_argument("--num_workers", type=int, default=8) parser.add_argument("--model_path", type=str, default="genscore.pth") parser.add_argument("--encoder", type=str, choices=["gt", "gatedgcn"], default="gt") parser.add_argument("--mode", type=str, choices=["lower", "higher"], default="lower") parser.add_argument("--finetune", action="store_true", default=False) parser.add_argument("--original_model_path", type=str, default=None) parser.add_argument("--lr", type=int, default=3) parser.add_argument("--weight_decay", type=int, default=5) parser.add_argument("--data_dir", type=str, required=True) parser.add_argument("--data_prefix", type=str, default="v2020_train") parser.add_argument("--valnum", type=int, default=1500) parser.add_argument("--seeds", type=int, default=126) parser.add_argument("--hidden_dim0", type=int, default=128) parser.add_argument("--hidden_dim", type=int, default=128) parser.add_argument("--n_gaussians", type=int, default=10) parser.add_argument("--dropout_rate", type=float, default=0.15) parser.add_argument("--dist_threhold", type=float, default=7.0) parser.add_argument("--dist_threhold2", type=float, default=5.0) return parser.parse_args() def _build_encoder(args): if args.encoder == "gt": ligmodel = GraphTransformer( in_channels=41, edge_features=10, num_hidden_channels=args.hidden_dim0, activ_fn=th.nn.SiLU(), transformer_residual=True, num_attention_heads=4, norm_to_apply="batch", dropout_rate=0.15, num_layers=6, ) protmodel = GraphTransformer( in_channels=41, edge_features=5, num_hidden_channels=args.hidden_dim0, activ_fn=th.nn.SiLU(), transformer_residual=True, num_attention_heads=4, norm_to_apply="batch", dropout_rate=0.15, num_layers=6, ) else: ligmodel = GatedGCN( in_channels=41, edge_features=10, num_hidden_channels=args.hidden_dim0, residual=True, dropout_rate=0.15, equivstable_pe=False, num_layers=6, ) protmodel = GatedGCN( in_channels=41, edge_features=5, num_hidden_channels=args.hidden_dim0, residual=True, dropout_rate=0.15, equivstable_pe=False, num_layers=6, ) return ligmodel, protmodel def main(): args = parse_args() args.device = "cuda" if th.cuda.is_available() else "cpu" data = PDBbindDataset( ids=f"{args.data_dir}/{args.data_prefix}_ids.npy", ligs=f"{args.data_dir}/{args.data_prefix}_lig.pt", prots=f"{args.data_dir}/{args.data_prefix}_prot.pt", ) train_inds, val_inds = data.train_and_test_split(valnum=args.valnum, seed=args.seeds) train_data = PDBbindDataset( ids=data.pdbids[train_inds], ligs=data.gls[train_inds], prots=data.gps[train_inds], labels=data.labels[train_inds], ) val_data = PDBbindDataset( ids=data.pdbids[val_inds], ligs=data.gls[val_inds], prots=data.gps[val_inds], labels=data.labels[val_inds], ) ligmodel, protmodel = _build_encoder(args) model = GenScore( ligmodel, protmodel, in_channels=args.hidden_dim0, hidden_dim=args.hidden_dim, n_gaussians=args.n_gaussians, dropout_rate=args.dropout_rate, dist_threhold=args.dist_threhold, ).to(args.device) if args.finetune: if args.original_model_path is None: raise ValueError('--original_model_path is required when --finetune is used.') checkpoint = th.load(args.original_model_path, map_location=th.device(args.device)) model.load_state_dict(checkpoint["model_state_dict"]) optimizer = th.optim.Adam( model.parameters(), lr=10**-args.lr, weight_decay=10**-args.weight_decay, ) train_loader = DataLoader( dataset=train_data, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, ) val_loader = DataLoader( dataset=val_data, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, ) stopper = EarlyStopping(patience=args.patience, mode=args.mode, filename=args.model_path) set_random_seed(args.seeds) for epoch in range(args.num_epochs): total_loss_train, mdn_loss_train, affi_loss_train, atom_loss_train, bond_loss_train = run_a_train_epoch( epoch, model, train_loader, optimizer, affi_weight=args.affi_weight, aux_weight=args.aux_weight, dist_threhold=args.dist_threhold2, device=args.device, ) if np.isinf(mdn_loss_train) or np.isnan(mdn_loss_train): print("Inf ERROR") break total_loss_val, mdn_loss_val, affi_loss_val, atom_loss_val, bond_loss_val = run_an_eval_epoch( model, val_loader, dist_threhold=args.dist_threhold2, affi_weight=args.affi_weight, aux_weight=args.aux_weight, device=args.device, ) early_stop = stopper.step(total_loss_val, model) print( "epoch {:d}/{:d}, total_loss_val {:.4f}, mdn_loss_val {:.4f}, " "affi_loss_val {:.4f}, atom_loss_val {:.4f}, bond_loss_val {:.4f}, " "best validation {:.4f}".format( epoch + 1, args.num_epochs, total_loss_val, mdn_loss_val, affi_loss_val, atom_loss_val, bond_loss_val, stopper.best_score, ) ) if early_stop: break stopper.load_checkpoint(model) train_metrics = run_an_eval_epoch( model, train_loader, dist_threhold=args.dist_threhold2, affi_weight=args.affi_weight, aux_weight=args.aux_weight, device=args.device, ) val_metrics = run_an_eval_epoch( model, val_loader, dist_threhold=args.dist_threhold2, affi_weight=args.affi_weight, aux_weight=args.aux_weight, device=args.device, ) print("train metrics:", train_metrics) print("validation metrics:", val_metrics) if __name__ == "__main__": main()