OneScience-Group/DiffDock
053
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 