Team Ai
Modelpublic

OneScience-Group/DiffDock

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes53downloads
sample_diffdock.py396 linesDownload Raw Back to scripts
1import sys2from pathlib import Path3 4DIR = Path(__file__).resolve().parent.parent5sys.path.insert(0, str(DIR))6import argparse7import copy8import csv9import os10from pathlib import Path11from types import SimpleNamespace12 13import numpy as np14import torch15import yaml16from rdkit.Chem import RemoveAllHs17from torch_geometric.loader import DataLoader18 19from onescience.datapipes.diffdock.process_mols import write_mol_with_coords20from onescience.utils.diffdock.diffusion_utils import get_t_schedule21from onescience.utils.diffdock.inference_utils import InferenceDataset, set_nones22from onescience.utils.diffdock.logging_utils import configure_logger, get_logger23from onescience.utils.diffdock.sampling import randomize_position, sampling24from onescience.utils.diffdock.validation import validate_sampling_entrypoint25 26try:27    from models.score_wrapper import load_model_args, load_score_model, model_uses_lm_embeddings28except ImportError:29    from models.score_wrapper import load_model_args, load_score_model, model_uses_lm_embeddings30 31 32def parse_args():33    parser = argparse.ArgumentParser()34    parser.add_argument("--config", required=True, help="Path to the sampling YAML config.")35    return parser.parse_args()36 37def _resolve_env_vars(obj):38    if isinstance(obj, str):39        return os.path.expandvars(obj)40    if isinstance(obj, dict):41        return {k: _resolve_env_vars(v) for k, v in obj.items()}42    if isinstance(obj, list):43        return [_resolve_env_vars(v) for v in obj]44    return obj45 46# def load_config(config_path):47#     with open(config_path, "r", encoding="utf-8") as handle:48#         return yaml.safe_load(handle) or {}49def load_config(config_path):50    with open(config_path, "r", encoding="utf-8") as handle:51        return _resolve_env_vars(yaml.safe_load(handle) or {})52 53 54def flatten_config(config):55    flat = {}56    for key, value in config.items():57        if isinstance(value, dict):58            flat.update(value)59        else:60            flat[key] = value61    return flat62 63 64def to_namespace(config):65    return SimpleNamespace(**config)66 67 68def resolve_device(device_name):69    if device_name in {None, "auto"}:70        return torch.device("cuda" if torch.cuda.is_available() else "cpu")71    return torch.device(device_name)72 73 74def load_inputs(args):75    if args.protein_ligand_csv is not None:76        with open(args.protein_ligand_csv, "r", encoding="utf-8", newline="") as handle:77            rows = list(csv.DictReader(handle))78        complex_names = set_nones([row.get("complex_name") for row in rows])79        protein_paths = set_nones([row.get("protein_path") for row in rows])80        protein_sequences = set_nones([row.get("protein_sequence") for row in rows])81        ligand_descriptions = set_nones([row.get("ligand_description") for row in rows])82    else:83        complex_names = [args.complex_name or "complex_0"]84        protein_paths = [args.protein_path]85        protein_sequences = [args.protein_sequence]86        ligand_descriptions = [args.ligand_description]87 88    complex_names = [name if name is not None else f"complex_{idx}" for idx, name in enumerate(complex_names)]89    return complex_names, protein_paths, protein_sequences, ligand_descriptions90 91 92def ensure_output_dirs(out_dir, complex_names):93    os.makedirs(out_dir, exist_ok=True)94    for name in complex_names:95        os.makedirs(os.path.join(out_dir, name), exist_ok=True)96 97 98def get_ligand_mol(complex_graph):99    mol = complex_graph.mol100    return mol[0] if isinstance(mol, (list, tuple)) else mol101 102 103def resolve_lm_embeddings_flag(args, model_args):104    lm_embeddings = getattr(args, "lm_embeddings", None)105    if lm_embeddings is None:106        return model_uses_lm_embeddings(model_args)107    return lm_embeddings108 109 110def build_inference_dataset(111    args,112    model_args,113    *,114    complex_names,115    protein_paths,116    protein_sequences,117    ligand_descriptions,118    lm_embeddings,119):120    return InferenceDataset(121        out_dir=args.out_dir,122        complex_names=complex_names,123        protein_files=protein_paths,124        ligand_descriptions=ligand_descriptions,125        protein_sequences=protein_sequences,126        lm_embeddings=lm_embeddings,127        receptor_radius=model_args.receptor_radius,128        remove_hs=model_args.remove_hs,129        c_alpha_max_neighbors=model_args.c_alpha_max_neighbors,130        all_atoms=model_args.all_atoms,131        atom_radius=model_args.atom_radius,132        atom_max_neighbors=model_args.atom_max_neighbors,133        knn_only_graph=not getattr(model_args, "not_knn_only_graph", False),134    )135 136 137def inference_graph_signature(model_args, lm_embeddings):138    return (139        getattr(model_args, "receptor_radius", None),140        getattr(model_args, "remove_hs", None),141        getattr(model_args, "c_alpha_max_neighbors", None),142        getattr(model_args, "all_atoms", None),143        getattr(model_args, "atom_radius", None),144        getattr(model_args, "atom_max_neighbors", None),145        not getattr(model_args, "not_knn_only_graph", False),146        bool(lm_embeddings),147    )148 149 150def extract_confidence_scores(confidence, confidence_model_args):151    confidence_scores = confidence152    if isinstance(getattr(confidence_model_args, "rmsd_classification_cutoff", None), list):153        confidence_scores = confidence_scores[:, 0]154    return np.asarray(confidence_scores.detach().cpu().numpy()).reshape(-1)155 156 157def load_optional_confidence_model(args, device, logger):158    confidence_model_dir = getattr(args, "confidence_model_dir", None)159    if confidence_model_dir is None:160        return None, None161 162    confidence_model_dir = Path(confidence_model_dir)163    confidence_checkpoint_path = confidence_model_dir / args.confidence_ckpt164    if not confidence_checkpoint_path.exists():165        raise FileNotFoundError(166            f"Confidence checkpoint not found: {confidence_checkpoint_path}. "167            "This example does not auto-download models."168        )169 170    confidence_model_args_preview = load_model_args(confidence_model_dir)171    validate_sampling_entrypoint(172        confidence_model_args_preview,173        context="DiffDock sampling confidence-model checkpoint",174        include_confidence=True,175        confidence_mode=True,176    )177 178    confidence_model, confidence_model_args, _ = load_score_model(179        model_dir=confidence_model_dir,180        ckpt=args.confidence_ckpt,181        device=device,182        no_parallel=True,183        confidence_mode=True,184        old=getattr(args, "old_confidence_model", False),185    )186    if hasattr(args, "crop_beyond"):187        confidence_model_args.crop_beyond = args.crop_beyond188    logger.info("Loaded optional confidence model from %s", confidence_model_dir)189    return confidence_model, confidence_model_args190 191 192def main():193    parsed = parse_args()194    raw_config = load_config(parsed.config)195    args = to_namespace(flatten_config(raw_config))196    device = resolve_device(getattr(args, "device", "auto"))197    validate_sampling_entrypoint(198        args,199        include_confidence=(200            getattr(args, "confidence_model_dir", None) is not None201            or getattr(args, "old_confidence_model", False)202        ),203    )204 205    configure_logger(getattr(args, "loglevel", "INFO"))206    logger = get_logger()207 208    model_dir = Path(args.model_dir)209    checkpoint_path = model_dir / args.ckpt210    if not checkpoint_path.exists():211        raise FileNotFoundError(212            f"Checkpoint not found: {checkpoint_path}. This example does not auto-download models."213        )214    score_model_args_preview = load_model_args(model_dir)215    validate_sampling_entrypoint(216        score_model_args_preview,217        context="DiffDock sampling score-model checkpoint",218    )219 220    model, score_model_args, t_to_sigma = load_score_model(221        model_dir=model_dir,222        ckpt=args.ckpt,223        device=device,224        no_parallel=True,225        old=getattr(args, "old_score_model", False),226    )227    if hasattr(args, "crop_beyond"):228        score_model_args.crop_beyond = args.crop_beyond229    logger.info("DiffDock sampling will run on %s", device)230 231    confidence_model, confidence_model_args = load_optional_confidence_model(args, device, logger)232 233    complex_names, protein_paths, protein_sequences, ligand_descriptions = load_inputs(args)234    ensure_output_dirs(args.out_dir, complex_names)235 236    score_lm_embeddings = resolve_lm_embeddings_flag(args, score_model_args)237    test_dataset = build_inference_dataset(238        args,239        score_model_args,240        complex_names=complex_names,241        protein_paths=protein_paths,242        protein_sequences=protein_sequences,243        ligand_descriptions=ligand_descriptions,244        lm_embeddings=score_lm_embeddings,245    )246    test_loader = DataLoader(dataset=test_dataset, batch_size=1, shuffle=False)247 248    confidence_loader = None249    if confidence_model is not None:250        confidence_lm_embeddings = resolve_lm_embeddings_flag(args, confidence_model_args)251        confidence_needs_independent_graph = (252            inference_graph_signature(score_model_args, score_lm_embeddings)253            != inference_graph_signature(confidence_model_args, confidence_lm_embeddings)254        )255        if confidence_needs_independent_graph:256            logger.info(257                "Confidence rerank requires independent inference graphs; building a separate confidence dataset."258            )259            confidence_dataset = build_inference_dataset(260                args,261                confidence_model_args,262                complex_names=complex_names,263                protein_paths=protein_paths,264                protein_sequences=protein_sequences,265                ligand_descriptions=ligand_descriptions,266                lm_embeddings=confidence_lm_embeddings,267            )268            confidence_loader = DataLoader(dataset=confidence_dataset, batch_size=1, shuffle=False)269        else:270            logger.info("Confidence rerank will reuse the score-model inference graphs.")271 272    tr_schedule = get_t_schedule(273        sigma_schedule=args.sigma_schedule,274        inference_steps=args.inference_steps,275        inf_sched_alpha=args.inf_sched_alpha,276        inf_sched_beta=args.inf_sched_beta,277    )278 279    failures = 0280    skipped = 0281    num_samples = args.samples_per_complex282    test_ds_size = len(test_dataset)283    logger.info("Size of test dataset: %s", test_ds_size)284 285    if confidence_loader is None:286        loader_iter = ((orig_complex_graph, None) for orig_complex_graph in test_loader)287    else:288        loader_iter = zip(test_loader, confidence_loader)289 290    for idx, (orig_complex_graph, confidence_orig_complex_graph) in enumerate(loader_iter):291        if not orig_complex_graph.success[0]:292            skipped += 1293            logger.warning(294                "Skipping %s because preprocessing failed.",295                test_dataset.complex_names[idx],296            )297            continue298        if confidence_orig_complex_graph is not None and not confidence_orig_complex_graph.success[0]:299            skipped += 1300            logger.warning(301                "Skipping %s because confidence preprocessing failed.",302                test_dataset.complex_names[idx],303            )304            continue305 306        try:307            data_list = [copy.deepcopy(orig_complex_graph) for _ in range(num_samples)]308            confidence_data_list = None309            if confidence_orig_complex_graph is not None:310                confidence_data_list = [copy.deepcopy(confidence_orig_complex_graph) for _ in range(num_samples)]311            randomize_position(312                data_list,313                score_model_args.no_torsion,314                args.no_random,315                score_model_args.tr_sigma_max,316                initial_noise_std_proportion=args.initial_noise_std_proportion,317                choose_residue=args.choose_residue,318            )319 320            ligand = get_ligand_mol(orig_complex_graph)321            data_list, confidence = sampling(322                data_list=data_list,323                model=model,324                inference_steps=args.actual_steps if args.actual_steps is not None else args.inference_steps,325                tr_schedule=tr_schedule,326                rot_schedule=tr_schedule,327                tor_schedule=tr_schedule,328                device=device,329                t_to_sigma=t_to_sigma,330                model_args=score_model_args,331                no_random=args.no_random,332                ode=args.ode,333                confidence_model=confidence_model,334                confidence_data_list=confidence_data_list,335                confidence_model_args=confidence_model_args,336                batch_size=args.batch_size,337                no_final_step_noise=args.no_final_step_noise,338                temp_sampling=[339                    args.temp_sampling_tr,340                    args.temp_sampling_rot,341                    args.temp_sampling_tor,342                ],343                temp_psi=[344                    args.temp_psi_tr,345                    args.temp_psi_rot,346                    args.temp_psi_tor,347                ],348                temp_sigma_data=[349                    args.temp_sigma_data_tr,350                    args.temp_sigma_data_rot,351                    args.temp_sigma_data_tor,352                ],353            )354 355            ligand_positions = np.asarray(356                [357                    complex_graph["ligand"].pos.cpu().numpy() + orig_complex_graph.original_center.cpu().numpy()358                    for complex_graph in data_list359                ]360            )361 362            rerank_order = np.arange(len(ligand_positions))363            confidence_scores = None364            if confidence is not None:365                confidence_scores = extract_confidence_scores(confidence, confidence_model_args)366                confidence_scores = np.nan_to_num(confidence_scores, nan=-1e-6)367                rerank_order = np.argsort(confidence_scores)[::-1]368                logger.info(369                    "Applied confidence rerank for %s. Ranked confidences: %s",370                    complex_names[idx],371                    np.array2string(confidence_scores[rerank_order], precision=4),372                )373 374            write_dir = os.path.join(args.out_dir, complex_names[idx])375            for rank, sample_idx in enumerate(rerank_order, start=1):376                pos = ligand_positions[sample_idx]377                mol_pred = copy.deepcopy(ligand)378                if score_model_args.remove_hs:379                    mol_pred = RemoveAllHs(mol_pred)380                filename = f"rank{rank}.sdf"381                if confidence_scores is not None:382                    filename = f"rank{rank}_conf{confidence_scores[sample_idx]:.4f}.sdf"383                write_mol_with_coords(mol_pred, pos, os.path.join(write_dir, filename))384 385        except Exception as exc:386            logger.warning("Failed on %s with error: %s", complex_names[idx], exc)387            failures += 1388 389    logger.info("Failed for %s / %s complexes.", failures, test_ds_size)390    logger.info("Skipped %s / %s complexes.", skipped, test_ds_size)391    logger.info("Results saved in %s", args.out_dir)392 393 394if __name__ == "__main__":395    main()396