DiffDock / scripts /sample_diffdock.py
OneScience's picture
Upload folder using huggingface_hub
c2767f4 verified
Raw
History Blame Contribute Delete
15.1 kB
import sys
from pathlib import Path
DIR = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(DIR))
import argparse
import copy
import csv
import os
from pathlib import Path
from types import SimpleNamespace
import numpy as np
import torch
import yaml
from rdkit.Chem import RemoveAllHs
from torch_geometric.loader import DataLoader
from onescience.datapipes.diffdock.process_mols import write_mol_with_coords
from onescience.utils.diffdock.diffusion_utils import get_t_schedule
from onescience.utils.diffdock.inference_utils import InferenceDataset, set_nones
from onescience.utils.diffdock.logging_utils import configure_logger, get_logger
from onescience.utils.diffdock.sampling import randomize_position, sampling
from onescience.utils.diffdock.validation import validate_sampling_entrypoint
try:
from models.score_wrapper import load_model_args, load_score_model, model_uses_lm_embeddings
except ImportError:
from models.score_wrapper import load_model_args, load_score_model, model_uses_lm_embeddings
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True, help="Path to the sampling YAML config.")
return parser.parse_args()
def _resolve_env_vars(obj):
if isinstance(obj, str):
return os.path.expandvars(obj)
if isinstance(obj, dict):
return {k: _resolve_env_vars(v) for k, v in obj.items()}
if isinstance(obj, list):
return [_resolve_env_vars(v) for v in obj]
return obj
# def load_config(config_path):
# with open(config_path, "r", encoding="utf-8") as handle:
# return yaml.safe_load(handle) or {}
def load_config(config_path):
with open(config_path, "r", encoding="utf-8") as handle:
return _resolve_env_vars(yaml.safe_load(handle) or {})
def flatten_config(config):
flat = {}
for key, value in config.items():
if isinstance(value, dict):
flat.update(value)
else:
flat[key] = value
return flat
def to_namespace(config):
return SimpleNamespace(**config)
def resolve_device(device_name):
if device_name in {None, "auto"}:
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
return torch.device(device_name)
def load_inputs(args):
if args.protein_ligand_csv is not None:
with open(args.protein_ligand_csv, "r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
complex_names = set_nones([row.get("complex_name") for row in rows])
protein_paths = set_nones([row.get("protein_path") for row in rows])
protein_sequences = set_nones([row.get("protein_sequence") for row in rows])
ligand_descriptions = set_nones([row.get("ligand_description") for row in rows])
else:
complex_names = [args.complex_name or "complex_0"]
protein_paths = [args.protein_path]
protein_sequences = [args.protein_sequence]
ligand_descriptions = [args.ligand_description]
complex_names = [name if name is not None else f"complex_{idx}" for idx, name in enumerate(complex_names)]
return complex_names, protein_paths, protein_sequences, ligand_descriptions
def ensure_output_dirs(out_dir, complex_names):
os.makedirs(out_dir, exist_ok=True)
for name in complex_names:
os.makedirs(os.path.join(out_dir, name), exist_ok=True)
def get_ligand_mol(complex_graph):
mol = complex_graph.mol
return mol[0] if isinstance(mol, (list, tuple)) else mol
def resolve_lm_embeddings_flag(args, model_args):
lm_embeddings = getattr(args, "lm_embeddings", None)
if lm_embeddings is None:
return model_uses_lm_embeddings(model_args)
return lm_embeddings
def build_inference_dataset(
args,
model_args,
*,
complex_names,
protein_paths,
protein_sequences,
ligand_descriptions,
lm_embeddings,
):
return InferenceDataset(
out_dir=args.out_dir,
complex_names=complex_names,
protein_files=protein_paths,
ligand_descriptions=ligand_descriptions,
protein_sequences=protein_sequences,
lm_embeddings=lm_embeddings,
receptor_radius=model_args.receptor_radius,
remove_hs=model_args.remove_hs,
c_alpha_max_neighbors=model_args.c_alpha_max_neighbors,
all_atoms=model_args.all_atoms,
atom_radius=model_args.atom_radius,
atom_max_neighbors=model_args.atom_max_neighbors,
knn_only_graph=not getattr(model_args, "not_knn_only_graph", False),
)
def inference_graph_signature(model_args, lm_embeddings):
return (
getattr(model_args, "receptor_radius", None),
getattr(model_args, "remove_hs", None),
getattr(model_args, "c_alpha_max_neighbors", None),
getattr(model_args, "all_atoms", None),
getattr(model_args, "atom_radius", None),
getattr(model_args, "atom_max_neighbors", None),
not getattr(model_args, "not_knn_only_graph", False),
bool(lm_embeddings),
)
def extract_confidence_scores(confidence, confidence_model_args):
confidence_scores = confidence
if isinstance(getattr(confidence_model_args, "rmsd_classification_cutoff", None), list):
confidence_scores = confidence_scores[:, 0]
return np.asarray(confidence_scores.detach().cpu().numpy()).reshape(-1)
def load_optional_confidence_model(args, device, logger):
confidence_model_dir = getattr(args, "confidence_model_dir", None)
if confidence_model_dir is None:
return None, None
confidence_model_dir = Path(confidence_model_dir)
confidence_checkpoint_path = confidence_model_dir / args.confidence_ckpt
if not confidence_checkpoint_path.exists():
raise FileNotFoundError(
f"Confidence checkpoint not found: {confidence_checkpoint_path}. "
"This example does not auto-download models."
)
confidence_model_args_preview = load_model_args(confidence_model_dir)
validate_sampling_entrypoint(
confidence_model_args_preview,
context="DiffDock sampling confidence-model checkpoint",
include_confidence=True,
confidence_mode=True,
)
confidence_model, confidence_model_args, _ = load_score_model(
model_dir=confidence_model_dir,
ckpt=args.confidence_ckpt,
device=device,
no_parallel=True,
confidence_mode=True,
old=getattr(args, "old_confidence_model", False),
)
if hasattr(args, "crop_beyond"):
confidence_model_args.crop_beyond = args.crop_beyond
logger.info("Loaded optional confidence model from %s", confidence_model_dir)
return confidence_model, confidence_model_args
def main():
parsed = parse_args()
raw_config = load_config(parsed.config)
args = to_namespace(flatten_config(raw_config))
device = resolve_device(getattr(args, "device", "auto"))
validate_sampling_entrypoint(
args,
include_confidence=(
getattr(args, "confidence_model_dir", None) is not None
or getattr(args, "old_confidence_model", False)
),
)
configure_logger(getattr(args, "loglevel", "INFO"))
logger = get_logger()
model_dir = Path(args.model_dir)
checkpoint_path = model_dir / args.ckpt
if not checkpoint_path.exists():
raise FileNotFoundError(
f"Checkpoint not found: {checkpoint_path}. This example does not auto-download models."
)
score_model_args_preview = load_model_args(model_dir)
validate_sampling_entrypoint(
score_model_args_preview,
context="DiffDock sampling score-model checkpoint",
)
model, score_model_args, t_to_sigma = load_score_model(
model_dir=model_dir,
ckpt=args.ckpt,
device=device,
no_parallel=True,
old=getattr(args, "old_score_model", False),
)
if hasattr(args, "crop_beyond"):
score_model_args.crop_beyond = args.crop_beyond
logger.info("DiffDock sampling will run on %s", device)
confidence_model, confidence_model_args = load_optional_confidence_model(args, device, logger)
complex_names, protein_paths, protein_sequences, ligand_descriptions = load_inputs(args)
ensure_output_dirs(args.out_dir, complex_names)
score_lm_embeddings = resolve_lm_embeddings_flag(args, score_model_args)
test_dataset = build_inference_dataset(
args,
score_model_args,
complex_names=complex_names,
protein_paths=protein_paths,
protein_sequences=protein_sequences,
ligand_descriptions=ligand_descriptions,
lm_embeddings=score_lm_embeddings,
)
test_loader = DataLoader(dataset=test_dataset, batch_size=1, shuffle=False)
confidence_loader = None
if confidence_model is not None:
confidence_lm_embeddings = resolve_lm_embeddings_flag(args, confidence_model_args)
confidence_needs_independent_graph = (
inference_graph_signature(score_model_args, score_lm_embeddings)
!= inference_graph_signature(confidence_model_args, confidence_lm_embeddings)
)
if confidence_needs_independent_graph:
logger.info(
"Confidence rerank requires independent inference graphs; building a separate confidence dataset."
)
confidence_dataset = build_inference_dataset(
args,
confidence_model_args,
complex_names=complex_names,
protein_paths=protein_paths,
protein_sequences=protein_sequences,
ligand_descriptions=ligand_descriptions,
lm_embeddings=confidence_lm_embeddings,
)
confidence_loader = DataLoader(dataset=confidence_dataset, batch_size=1, shuffle=False)
else:
logger.info("Confidence rerank will reuse the score-model inference graphs.")
tr_schedule = get_t_schedule(
sigma_schedule=args.sigma_schedule,
inference_steps=args.inference_steps,
inf_sched_alpha=args.inf_sched_alpha,
inf_sched_beta=args.inf_sched_beta,
)
failures = 0
skipped = 0
num_samples = args.samples_per_complex
test_ds_size = len(test_dataset)
logger.info("Size of test dataset: %s", test_ds_size)
if confidence_loader is None:
loader_iter = ((orig_complex_graph, None) for orig_complex_graph in test_loader)
else:
loader_iter = zip(test_loader, confidence_loader)
for idx, (orig_complex_graph, confidence_orig_complex_graph) in enumerate(loader_iter):
if not orig_complex_graph.success[0]:
skipped += 1
logger.warning(
"Skipping %s because preprocessing failed.",
test_dataset.complex_names[idx],
)
continue
if confidence_orig_complex_graph is not None and not confidence_orig_complex_graph.success[0]:
skipped += 1
logger.warning(
"Skipping %s because confidence preprocessing failed.",
test_dataset.complex_names[idx],
)
continue
try:
data_list = [copy.deepcopy(orig_complex_graph) for _ in range(num_samples)]
confidence_data_list = None
if confidence_orig_complex_graph is not None:
confidence_data_list = [copy.deepcopy(confidence_orig_complex_graph) for _ in range(num_samples)]
randomize_position(
data_list,
score_model_args.no_torsion,
args.no_random,
score_model_args.tr_sigma_max,
initial_noise_std_proportion=args.initial_noise_std_proportion,
choose_residue=args.choose_residue,
)
ligand = get_ligand_mol(orig_complex_graph)
data_list, confidence = sampling(
data_list=data_list,
model=model,
inference_steps=args.actual_steps if args.actual_steps is not None else args.inference_steps,
tr_schedule=tr_schedule,
rot_schedule=tr_schedule,
tor_schedule=tr_schedule,
device=device,
t_to_sigma=t_to_sigma,
model_args=score_model_args,
no_random=args.no_random,
ode=args.ode,
confidence_model=confidence_model,
confidence_data_list=confidence_data_list,
confidence_model_args=confidence_model_args,
batch_size=args.batch_size,
no_final_step_noise=args.no_final_step_noise,
temp_sampling=[
args.temp_sampling_tr,
args.temp_sampling_rot,
args.temp_sampling_tor,
],
temp_psi=[
args.temp_psi_tr,
args.temp_psi_rot,
args.temp_psi_tor,
],
temp_sigma_data=[
args.temp_sigma_data_tr,
args.temp_sigma_data_rot,
args.temp_sigma_data_tor,
],
)
ligand_positions = np.asarray(
[
complex_graph["ligand"].pos.cpu().numpy() + orig_complex_graph.original_center.cpu().numpy()
for complex_graph in data_list
]
)
rerank_order = np.arange(len(ligand_positions))
confidence_scores = None
if confidence is not None:
confidence_scores = extract_confidence_scores(confidence, confidence_model_args)
confidence_scores = np.nan_to_num(confidence_scores, nan=-1e-6)
rerank_order = np.argsort(confidence_scores)[::-1]
logger.info(
"Applied confidence rerank for %s. Ranked confidences: %s",
complex_names[idx],
np.array2string(confidence_scores[rerank_order], precision=4),
)
write_dir = os.path.join(args.out_dir, complex_names[idx])
for rank, sample_idx in enumerate(rerank_order, start=1):
pos = ligand_positions[sample_idx]
mol_pred = copy.deepcopy(ligand)
if score_model_args.remove_hs:
mol_pred = RemoveAllHs(mol_pred)
filename = f"rank{rank}.sdf"
if confidence_scores is not None:
filename = f"rank{rank}_conf{confidence_scores[sample_idx]:.4f}.sdf"
write_mol_with_coords(mol_pred, pos, os.path.join(write_dir, filename))
except Exception as exc:
logger.warning("Failed on %s with error: %s", complex_names[idx], exc)
failures += 1
logger.info("Failed for %s / %s complexes.", failures, test_ds_size)
logger.info("Skipped %s / %s complexes.", skipped, test_ds_size)
logger.info("Results saved in %s", args.out_dir)
if __name__ == "__main__":
main()