GenScore / models /train.py
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
9d6a2a3 verified
Raw
History Blame Contribute Delete
7.5 kB
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()