Download benchmarks/_common.py from OneScience-Group/GenScore: direct link, hf CLI and curl.
- Browser
- Download file 3.27 kB
-
https://huggingface.co/OneScience-Group/GenScore/resolve/main/benchmarks/_common.py
- Command line
-
hf download hf://OneScience-Group/GenScore/benchmarks/_common.py
-
curl -L -o _common.py https://huggingface.co/OneScience-Group/GenScore/resolve/main/benchmarks/_common.py
3.27 kB
| import sys | |
| from pathlib import Path | |
| _DIR = Path(__file__).resolve().parent.parent | |
| sys.path.insert(0, str(_DIR)) | |
| import argparse | |
| import torch as th | |
| import torch.multiprocessing | |
| from torch_geometric.loader import DataLoader | |
| from onescience.datapipes.genscore.data import PDBbindDataset | |
| from onescience.metrics.genscore.utils import run_an_eval_epoch | |
| from models.inference import _build_encoder, scoring | |
| from models.model.model import GenScore | |
| torch.multiprocessing.set_sharing_strategy("file_system") | |
| def add_model_args(parser): | |
| parser.add_argument("--model-path", required=True, help="Path to a trained GenScore checkpoint.") | |
| parser.add_argument("--encoder", choices=["gt", "gatedgcn"], default="gatedgcn") | |
| parser.add_argument("--batch-size", type=int, default=128) | |
| parser.add_argument("--num-workers", type=int, default=10) | |
| parser.add_argument("--cutoff", type=float, default=10.0) | |
| parser.add_argument("--outprefix", default="gatedgcn1x5") | |
| parser.add_argument("--dist-threhold", type=float, default=5.0) | |
| 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) | |
| def runtime_kwargs(args): | |
| return { | |
| "batch_size": args.batch_size, | |
| "dist_threhold": args.dist_threhold, | |
| "device": "cuda" if th.cuda.is_available() else "cpu", | |
| "num_workers": args.num_workers, | |
| "num_node_featsp": 41, | |
| "num_node_featsl": 41, | |
| "num_edge_featsp": 5, | |
| "num_edge_featsl": 10, | |
| "hidden_dim0": args.hidden_dim0, | |
| "hidden_dim": args.hidden_dim, | |
| "n_gaussians": args.n_gaussians, | |
| "dropout_rate": args.dropout_rate, | |
| } | |
| def score_ligand_file(prot, lig, args, parallel=False): | |
| return scoring( | |
| prot=prot, | |
| lig=lig, | |
| modpath=args.model_path, | |
| cut=args.cutoff, | |
| gen_pocket=False, | |
| reflig=None, | |
| encoder=args.encoder, | |
| explicit_H=False, | |
| use_chirality=True, | |
| parallel=parallel, | |
| **runtime_kwargs(args), | |
| ) | |
| def score_preprocessed(ids, prots, ligs, args): | |
| kwargs = runtime_kwargs(args) | |
| data = PDBbindDataset(ids=ids, prots=prots, ligs=ligs) | |
| loader = DataLoader( | |
| dataset=data, | |
| batch_size=kwargs["batch_size"], | |
| shuffle=False, | |
| num_workers=kwargs["num_workers"], | |
| ) | |
| ligmodel, protmodel = _build_encoder(args.encoder, kwargs) | |
| model = GenScore( | |
| ligmodel, | |
| protmodel, | |
| in_channels=kwargs["hidden_dim0"], | |
| hidden_dim=kwargs["hidden_dim"], | |
| n_gaussians=kwargs["n_gaussians"], | |
| dropout_rate=kwargs["dropout_rate"], | |
| dist_threhold=kwargs["dist_threhold"], | |
| ).to(kwargs["device"]) | |
| checkpoint = th.load(args.model_path, map_location=th.device(kwargs["device"])) | |
| model.load_state_dict(checkpoint["model_state_dict"]) | |
| preds = run_an_eval_epoch( | |
| model, | |
| loader, | |
| pred=True, | |
| dist_threhold=kwargs["dist_threhold"], | |
| device=kwargs["device"], | |
| ) | |
| return data.pdbids, preds | |
| def formatter(): | |
| return argparse.ArgumentDefaultsHelpFormatter | |