dl_binder_design / scripts /inference.py
wuxing0105's picture
Upload folder using huggingface_hub
9ae74ae verified
Raw
History Blame Contribute Delete
4.99 kB
#!/usr/bin/env python3
"""Unified PDB-directory inference entry point for dl_binder_design."""
from __future__ import annotations
import argparse
import json
import subprocess
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
CONFIG_PATH = ROOT / "conf" / "config.json"
MPNN_SCRIPT = ROOT / "model" / "mpnn_fr" / "dl_interface_design.py"
AF2_SCRIPT = ROOT / "model" / "af2_initial_guess" / "predict.py"
def load_config() -> dict:
with CONFIG_PATH.open(encoding="utf-8") as handle:
return json.load(handle)
def package_path(value: str) -> Path:
path = Path(value).expanduser()
return path if path.is_absolute() else ROOT / path
def require_file(path: Path, label: str) -> None:
if not path.is_file() or path.stat().st_size == 0:
raise SystemExit(f"Missing {label}: {path}")
def run(command: list[str]) -> None:
print("Running:", " ".join(command), flush=True)
subprocess.run(command, cwd=ROOT, check=True)
def run_mpnn(args: argparse.Namespace, config: dict, output_dir: Path) -> Path:
checkpoint = package_path(args.mpnn_checkpoint)
require_file(checkpoint, "ProteinMPNN checkpoint")
require_file(
ROOT / "model" / "mpnn_fr" / "ProteinMPNN" / "protein_mpnn_utils.py",
"ProteinMPNN source file",
)
output_dir.mkdir(parents=True, exist_ok=True)
command = [
sys.executable,
str(MPNN_SCRIPT),
"-pdbdir",
str(package_path(args.input_dir)),
"-outpdbdir",
str(output_dir),
"-checkpoint_name",
str(output_dir.parent / f"{output_dir.name}.checkpoint"),
"-checkpoint_path",
str(checkpoint),
"-relax_cycles",
str(args.relax_cycles),
"-seqs_per_struct",
str(args.seqs_per_struct),
]
if args.runlist:
command.extend(["-runlist", str(package_path(args.runlist))])
if args.debug:
command.append("-debug")
run(command)
return output_dir
def run_af2(args: argparse.Namespace, input_dir: Path, output_dir: Path) -> None:
require_file(
ROOT / "weight" / "AlphaFold2" / "params" / "params_model_1_ptm.npz",
"AlphaFold2 model_1_ptm parameters",
)
output_dir.mkdir(parents=True, exist_ok=True)
command = [
sys.executable,
str(AF2_SCRIPT),
"-pdbdir",
str(input_dir),
"-outpdbdir",
str(output_dir),
"-checkpoint_name",
str(output_dir.parent / f"{output_dir.name}.checkpoint"),
"-scorefilename",
str(output_dir.parent / f"{output_dir.name}.sc"),
"-recycle",
str(args.recycle),
]
if args.runlist:
command.extend(["-runlist", str(package_path(args.runlist))])
if args.debug:
command.append("-debug")
run(command)
def add_shared_arguments(parser: argparse.ArgumentParser, config: dict) -> None:
parser.add_argument("--input-dir", default=config["sample_input_dir"])
parser.add_argument("--output-dir", default=config["output_dir"])
parser.add_argument("--runlist", default="")
parser.add_argument("--debug", action="store_true")
def main() -> None:
config = load_config()
parser = argparse.ArgumentParser(description=__doc__)
subparsers = parser.add_subparsers(dest="stage", required=True)
mpnn = subparsers.add_parser("mpnn", help="run ProteinMPNN with optional FastRelax")
add_shared_arguments(mpnn, config)
mpnn.add_argument("--mpnn-checkpoint", default=config["proteinmpnn_checkpoint"])
mpnn.add_argument("--relax-cycles", type=int, default=config["relax_cycles"])
mpnn.add_argument("--seqs-per-struct", type=int, default=config["seqs_per_struct"])
af2 = subparsers.add_parser("af2", help="run AlphaFold2 initial-guess prediction")
add_shared_arguments(af2, config)
af2.add_argument("--recycle", type=int, default=config["recycle"])
pipeline = subparsers.add_parser("pipeline", help="run ProteinMPNN/FastRelax then AF2")
add_shared_arguments(pipeline, config)
pipeline.add_argument("--mpnn-checkpoint", default=config["proteinmpnn_checkpoint"])
pipeline.add_argument("--relax-cycles", type=int, default=config["relax_cycles"])
pipeline.add_argument("--seqs-per-struct", type=int, default=config["seqs_per_struct"])
pipeline.add_argument("--recycle", type=int, default=config["recycle"])
args = parser.parse_args()
base_output = package_path(args.output_dir)
if args.stage == "mpnn":
run_mpnn(args, config, base_output)
elif args.stage == "af2":
run_af2(args, package_path(args.input_dir), base_output)
else:
mpnn_output = run_mpnn(args, config, base_output / "mpnn")
# ProteinMPNN changes output tags, so an input-stage runlist cannot be
# reused for the AF2 stage. AF2 processes all structures just produced.
args.runlist = ""
run_af2(args, mpnn_output, base_output / "af2")
if __name__ == "__main__":
main()