OneScience-Group/SurfDock
025
1"""2caoduanhua : we should to implemented a parapllel version of evaluate.py for a large dataset3"""4 5import copy6import os7import sys8 9SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))10PROJECT_DIR = os.path.dirname(SCRIPT_DIR)11MODEL_DIR = os.path.join(PROJECT_DIR, "model")12if MODEL_DIR not in sys.path:13 sys.path.insert(0, MODEL_DIR)14 15import torch16import time17from argparse import ArgumentParser, Namespace, FileType18from datetime import datetime19from functools import partial20import numpy as np21import wandb22from biopandas.pdb import PandasPdb23from rdkit import RDLogger24from rdkit.Chem import RemoveHs,AllChem25from datasets.process_mols import write_mol_with_coords, generate_conformer26from torch_geometric.loader import DataLoader27from datasets.pdbbind import PDBBind, read_mol28from utils.diffusion_utils import t_to_sigma as t_to_sigma_compl, get_t_schedule29from utils.sampling import randomize_position, sampling30from utils.utils import get_model, get_symmetry_rmsd, remove_all_hs, read_strings_from_txt, ExponentialMovingAverage31from utils.visualise import PDBFile32from tqdm import tqdm33from loguru import logger34torch.multiprocessing.set_sharing_strategy('file_system')35RDLogger.DisableLog('rdApp.*')36import yaml37 38cache_name = datetime.now().strftime('date%d-%m_time%H-%M-%S.%f')39parser = ArgumentParser()40parser.add_argument('--config', type=FileType(mode='r'), default=None)41parser.add_argument('--model_dir', type=str, default=None, help='Path to folder with trained score model and hyperparameters')42parser.add_argument('--ckpt', type=str, default=None, help='Checkpoint to use inside the folder')43parser.add_argument('--confidence_model_dir', type=str, default=None, help='Path to folder with trained confidence model and hyperparameters')44parser.add_argument('--confidence_ckpt', type=str, default=None, help='Checkpoint to use inside the folder')45parser.add_argument('--model_version', type=str, default='version3', help='version of mdn model')46# save docking result or not47parser.add_argument('--save_docking_result', action='store_true', default=False, help='Whether to save docking result')48# put ligand to pocket center49parser.add_argument('--ligand_to_pocket_center', action='store_true', default=False, help='Whether to put ligand on pocket center')50# use_noise_to_rank51parser.add_argument('--use_noise_to_rank', action='store_true', default=False, help='Whether to run the probability flow ODE')52parser.add_argument('--num_cpu', type=int, default=None, help='if this is a number instead of none, the max number of cpus used by torch will be set to this.')53parser.add_argument('--run_name', type=str, default='test_ns_48_nv_10_layer_62023-06-25_07-54-08_model', help='')54parser.add_argument('--project', type=str, default='ligbind_inf_test_mdn', help='')55parser.add_argument('--surface_path', type=str, default='~/PDBBind_processed_8A_surface/', help='test dataset surface path')56parser.add_argument('--esm_embeddings_path', type=str, default='~/DeepLearningForDock/datasets/equibind_and_diffdock_dataset/PDBBIND/esm_embedding/esm_embedding_pocket_for_train/esm2_3billion_embeddings.pt', help='test dataset esmbedding path')57parser.add_argument('--out_dir', type=str, default='~/diffScreen/test_workdir/mdn_result_40', help='Where to save results to')58parser.add_argument('--batch_size', type=int, default=40, help='Number of poses to sample in parallel')59parser.add_argument('--cache_path', type=str, default='~/DeepLearningForDock/datasets/equibind_and_diffdock_dataset/PDBBIND/cache_PDBBIND_pocket_8A', help='Folder from where to load/restore cached dataset')60parser.add_argument('--data_dir', type=str, default='~/DeepLearningForDock/datasets/equibind_and_diffdock_dataset/PDBBIND/PDBBind_pocket_8A/', help='Folder containing original structures')61parser.add_argument('--split_path', type=str, default='~/DeepLearningForDock/DiffDockForScreen/diffScreen/data/splits/timesplit_test', help='Path of file defining the split')62parser.add_argument('--no_overlap_names_path', type=str, default='~/DeepLearningForDock/DiffDockForScreen/diffScreen/data/splits/timesplit_test_no_rec_overlap', help='Path text file with the folder names in the test set that have no receptor overlap with the train set')63parser.add_argument('--no_model', action='store_true', default=False, help='Whether to return seed conformer without running model')64parser.add_argument('--no_random', action='store_true', default=False, help='Whether to add randomness in diffusion steps')65parser.add_argument('--no_final_step_noise', action='store_true', default=False, help='Whether to add noise after the final step')66parser.add_argument('--ode', action='store_true', default=False, help='Whether to run the probability flow ODE')67parser.add_argument('--wandb', action='store_true', default=True, help='')68parser.add_argument('--wandb_dir', type=str, default='~/diffScreen/test_workdir', help='Folder in which to save wandb logs')69parser.add_argument('--inference_steps', type=int, default=20, help='Number of denoising steps')70parser.add_argument('--multi_seed_conformer', action='store_true', default=False, help='Whether to use multi_seed_conformer in inference steps')71parser.add_argument('--limit_complexes', type=int, default=0, help='Limit to the number of complexes')72parser.add_argument('--num_workers', type=int, default=1, help='Number of workers for dataset creation')73parser.add_argument('--tqdm', action='store_true', default=False, help='Whether to show progress bar')74parser.add_argument('--save_visualisation', action='store_true', default=False, help='Whether to save visualizations')75parser.add_argument('--samples_per_complex', type=int, default=40, help='Number of poses to sample for each complex')76parser.add_argument('--actual_steps', type=int, default=None, help='')77parser.add_argument('--mdn_dist_threshold_test', type=float, default=None, help='mdn_dist_threshold_test')78# force_minimized param79parser.add_argument('--force_optimize', action='store_true', default=False, help='')80args = parser.parse_args()81def main_function():82 if accelerator.is_local_main_process:83 if args.wandb:84 wandb.login(key = 'yourkey')85 run = wandb.init(86 entity='SurfDock',87 settings=wandb.Settings(start_method="fork"),88 project=args.project,89 name=args.run_name,90 dir = args.wandb_dir,91 config=args92 )93 if args.config:94 config_dict = yaml.load(args.config, Loader=yaml.FullLoader)95 arg_dict = args.__dict__96 for key, value in config_dict.items():97 if isinstance(value, list):98 for v in value:99 arg_dict[key].append(v)100 else:101 arg_dict[key] = value102 if args.out_dir is None: args.out_dir = f'inference_out_dir_not_specified/{args.run_name}'103 os.makedirs(args.out_dir, exist_ok=True)104 with open(f'{args.model_dir}/model_parameters.yml') as f:105 score_model_args = Namespace(**yaml.full_load(f))106 107 if args.confidence_model_dir is not None:108 with open(f'{args.confidence_model_dir}/model_parameters.yml') as f:109 confidence_args = Namespace(**yaml.full_load(f))110 # 111 confidence_args.transfer_weights = False112 confidence_args.use_original_model_cache = True113 confidence_args.original_model_dir = None114 confidence_args.mdn_dist_threshold_test = args.mdn_dist_threshold_test if args.mdn_dist_threshold_test is not None else 5.0115 if not hasattr(confidence_args,'mdn_dist_threshold_train'):116 confidence_args.mdn_dist_threshold_train =7.0117 118 if args.force_optimize:119 logger.info('Using ForceField for energy minimized!')120 test_dataset = PDBBind(transform=None, root=args.data_dir, limit_complexes=args.limit_complexes,121 receptor_radius=score_model_args.receptor_radius,122 cache_path=args.cache_path, split_path=args.split_path,123 remove_hs=score_model_args.remove_hs, max_lig_size=None,124 c_alpha_max_neighbors=score_model_args.c_alpha_max_neighbors,125 matching=not score_model_args.no_torsion, keep_original=True,126 popsize=score_model_args.matching_popsize,127 maxiter=score_model_args.matching_maxiter,128 all_atoms=score_model_args.all_atoms,129 atom_radius=score_model_args.atom_radius,130 atom_max_neighbors=score_model_args.atom_max_neighbors,131 esm_embeddings_path=args.esm_embeddings_path,132 require_ligand=True,133 num_workers=args.num_workers,surface_path = args.surface_path)134 135 test_loader = DataLoader(dataset=test_dataset, batch_size=1, shuffle=False)136 if args.confidence_model_dir is not None:137 if not (confidence_args.use_original_model_cache or confidence_args.transfer_weights):138 # if the confidence model uses the same type of data as the original model then we do not need this dataset and can just use the complexes139 logger.info('HAPPENING | confidence model uses different type of graphs than the score model. Loading (or creating if not existing) the data for the confidence model now.')140 confidence_test_dataset = PDBBind(transform=None, root=args.data_dir, limit_complexes=args.limit_complexes,141 receptor_radius=confidence_args.receptor_radius,142 cache_path=args.cache_path, split_path=args.split_path,143 remove_hs=confidence_args.remove_hs, max_lig_size=None, c_alpha_max_neighbors=confidence_args.c_alpha_max_neighbors,144 matching=not confidence_args.no_torsion, keep_original=True,145 popsize=confidence_args.matching_popsize,146 maxiter=confidence_args.matching_maxiter,147 all_atoms=confidence_args.all_atoms,148 atom_radius=confidence_args.atom_radius,149 atom_max_neighbors=confidence_args.atom_max_neighbors,150 esm_embeddings_path= args.esm_embeddings_path, require_ligand=True,151 num_workers=args.num_workers,surface_path = args.surface_path)152 confidence_complex_dict = {d.name: d for d in confidence_test_dataset}153 154 t_to_sigma = partial(t_to_sigma_compl, args=score_model_args)155 156 if not args.no_model:157 model = get_model(score_model_args, device, t_to_sigma=t_to_sigma, no_parallel=True,model_type = score_model_args.model_type)158 state_dict = torch.load(f'{args.model_dir}/{args.ckpt}', map_location=torch.device('cpu'))159 if args.ckpt == 'last_model.pt':160 model_state_dict = state_dict['model']161 ema_weights_state = state_dict['ema_weights']162 model.load_state_dict(model_state_dict, strict=True)163 ema_weights = ExponentialMovingAverage(model.parameters(), decay=score_model_args.ema_rate)164 ema_weights.load_state_dict(ema_weights_state, device=device)165 ema_weights.copy_to(model.parameters())166 else:167 model.load_state_dict(state_dict, strict=False)168 model = model.to(device)169 model.eval()170 logger.info('loaded model weight for score model')171 if args.confidence_model_dir is not None:172 if confidence_args.transfer_weights:173 with open(f'{confidence_args.original_model_dir}/model_parameters.yml') as f:174 confidence_model_args = Namespace(**yaml.full_load(f))175 else:176 confidence_model_args = confidence_args177 178 confidence_model = get_model(confidence_model_args, device, t_to_sigma=t_to_sigma, no_parallel=True,179 model_type = confidence_model_args.model_type)180 state_dict = torch.load(f'{args.confidence_model_dir}/{args.confidence_ckpt}', map_location=torch.device('cpu'))181 confidence_model.load_state_dict(state_dict, strict=True)182 confidence_model = confidence_model.to(device)183 confidence_model.eval()184 else:185 confidence_model = None186 confidence_args = None187 confidence_model_args = None188 189 tr_schedule = get_t_schedule(inference_steps=args.inference_steps)190 rot_schedule = tr_schedule191 tor_schedule = tr_schedule192 logger.info('t schedule', tr_schedule)193 194 rmsds_list, obrmsds, centroid_distances_list, failures, skipped, min_cross_distances_list, base_min_cross_distances_list, confidences_list, names_list = [], [], [], 0, 0, [], [], [], []195 run_times, min_self_distances_list, without_rec_overlap_list = [], [], []196 N = args.samples_per_complex197 names_no_rec_overlap = read_strings_from_txt(args.no_overlap_names_path)198 names_test_all = read_strings_from_txt(args.split_path)199 name_map_idx = {name: idx for idx, name in enumerate(names_test_all)}200 idx_map_name = {idx: name for idx, name in enumerate(names_test_all)}201 202 logger.info('Size of test dataset: ', len(test_dataset))203 204 model = accelerator.prepare(model)205 test_loader= accelerator.prepare(test_loader)206 confidence_model = accelerator.prepare(confidence_model)207 208 for idx, orig_complex_graph in tqdm(enumerate(test_loader),total = len(test_loader),disable= not accelerator.is_local_main_process):209 if confidence_model is not None and not (confidence_args.use_original_model_cache or210 confidence_args.transfer_weights) and orig_complex_graph.name[0] not in confidence_complex_dict.keys():211 skipped += 1212 logger.info(f"HAPPENING | The confidence dataset did not contain {orig_complex_graph.name[0]}. We are skipping this complex.")213 continue214 success = 0215 sample_count_failed = 0216 while not success: # keep trying in case of failure (sometimes stochastic)217 try:218 success = 1219 data_list = [copy.deepcopy(orig_complex_graph) for _ in range(N)]220 221 if args.multi_seed_conformer:222 # random multi seed conformers223 if test_dataset.require_ligand:224 for data in data_list:225 mol_rdkit = copy.deepcopy(data.mol[0])226 mol_rdkit.RemoveAllConformers()227 mol_rdkit = AllChem.AddHs(mol_rdkit)228 generate_conformer(mol_rdkit)229 mol_rdkit = RemoveHs(mol_rdkit, sanitize=True)230 data.mol = [mol_rdkit]231 randomize_position(data_list, score_model_args.no_torsion, args.no_random, score_model_args.tr_sigma_max,ligand_to_pocket_center = args.ligand_to_pocket_center)232 233 pdb = None234 if args.save_visualisation:235 visualization_list = []236 for idx, graph in enumerate(data_list):237 # raw pose238 lig = read_mol(args.data_dir, graph['name'][0], remove_hs=score_model_args.remove_hs)239 pdb = PDBFile(lig)240 pdb.add(lig, 0, 0)241 # pose rdkit matching242 pdb.add((orig_complex_graph['ligand'].pos + orig_complex_graph.original_center).detach().cpu(), 1, 0)243 # logger.info(orig_complex_graph['ligand'].pos.shape,orig_complex_graph.original_center.shape)244 # logger.info(graph['ligand'].pos.device,graph.original_center.device)245 # random rdkit matching246 pdb.add((graph['ligand'].pos + (graph.original_center).detach().cpu()), part=1, order=1)247 visualization_list.append(pdb)248 else:249 visualization_list = None250 251 rec_path = os.path.join(args.data_dir, data_list[0]["name"][0], f'{data_list[0]["name"][0]}_pocket.pdb')252 if not os.path.exists(rec_path):253 rec_path = os.path.join(args.data_dir, data_list[0]["name"][0], f'{data_list[0]["name"][0]}_protein_obabel_reduce.pdb')254 rec = PandasPdb().read_pdb(rec_path)255 rec_df = rec.df['ATOM']256 receptor_pos = rec_df[['x_coord', 'y_coord', 'z_coord']].to_numpy().squeeze().astype(257 np.float32) - orig_complex_graph.original_center.cpu().numpy()258 receptor_pos = np.tile(receptor_pos, (N, 1, 1))259 start_time = time.time()260 if not args.no_model:261 if confidence_model is not None and not (262 confidence_args.use_original_model_cache or confidence_args.transfer_weights):263 confidence_data_list = [copy.deepcopy(confidence_complex_dict[orig_complex_graph.name[0]]) for _ in264 range(N)]265 else:266 confidence_data_list = None267 268 data_list, confidence = sampling(data_list=data_list, model=model,269 inference_steps=args.actual_steps if args.actual_steps is not None else args.inference_steps,270 tr_schedule=tr_schedule, rot_schedule=rot_schedule,271 tor_schedule=tor_schedule,272 device=device, t_to_sigma=t_to_sigma, model_args=score_model_args,273 no_random=args.no_random,274 ode=args.ode, visualization_list=visualization_list,275 confidence_model=confidence_model,276 confidence_data_list=confidence_data_list,277 confidence_model_args=confidence_model_args,278 batch_size=args.batch_size,279 no_final_step_noise=args.no_final_step_noise,args = args)280 281 confidence = confidence.cpu().detach().numpy()282 283 run_times.append(time.time() - start_time)284 if score_model_args.no_torsion: orig_complex_graph['ligand'].orig_pos = (orig_complex_graph['ligand'].pos.cpu().numpy() + orig_complex_graph.original_center.cpu().numpy())285 286 filterHs = torch.not_equal(data_list[0]['ligand'].x[:, 0], 0).cpu().numpy()287 288 if isinstance(orig_complex_graph['ligand'].orig_pos, list):289 orig_complex_graph['ligand'].orig_pos = orig_complex_graph['ligand'].orig_pos[0]290 291 ligand_pos = np.asarray(292 [complex_graph['ligand'].pos.cpu().numpy()[filterHs] for complex_graph in data_list])293 orig_ligand_pos = np.expand_dims(294 orig_complex_graph['ligand'].orig_pos[filterHs],295 axis=0) # - orig_complex_graph.original_center.cpu().numpy() ,since get_idx have done this function!296 297 try:298 mol = remove_all_hs(orig_complex_graph.mol[0])299 rmsd = get_symmetry_rmsd(mol, orig_ligand_pos[0], [l for l in ligand_pos])300 except Exception as e:301 logger.info("Using non corrected RMSD because of the error", e)302 rmsd = np.sqrt(((ligand_pos - orig_ligand_pos) ** 2).sum(axis=2).mean(axis=1))303 rmsds_list.append(rmsd)304 centroid_distance = np.linalg.norm(ligand_pos.mean(axis=1) - orig_ligand_pos.mean(axis=1), axis=1)305 # if confidence is not None and isinstance(confidence_args.rmsd_classification_cutoff, list):306 # confidence = confidence[:, 0]307 if confidence is not None:308 confidence = np.array(confidence)#.cpu().numpy()309 re_order = np.argsort(confidence)[::-1]310 # logger.info(confidence, re_order,rmsd)311 logger.info(orig_complex_graph['name'], ' rmsd', np.around(rmsd, 1)[re_order], ' centroid distance',312 np.around(centroid_distance, 1)[re_order], ' confidences ', np.around(confidence, 4)[re_order])313 confidences_list.append(confidence)314 315 else:316 logger.info(orig_complex_graph['name'], ' rmsd', np.around(rmsd, 1), ' centroid distance',317 np.around(centroid_distance, 1))318 319 """ add a save command by caoduanhua to save the last state of ligand"""320 ########################################################################321 if args.save_docking_result:322 ligand_pos_add_center = np.asarray([complex_graph['ligand'].pos.cpu().numpy() + orig_complex_graph.original_center.cpu().numpy() for complex_graph in data_list])323 lig = orig_complex_graph.mol[0]324 # save predictions325 326 write_dir = f'{args.out_dir}/docking_result_{os.path.basename(args.split_path)}/{orig_complex_graph.name[0]}'327 os.makedirs(write_dir, exist_ok=True)328 for pos,rmsd_i,score in zip(ligand_pos_add_center,rmsd,confidence):329 mol_pred = copy.deepcopy(lig)330 if score_model_args.remove_hs: mol_pred = RemoveHs(mol_pred)331 # if rank == 0: write_mol_with_coords(mol_pred, pos, os.path.join(write_dir, f'rank{rank+1}.sdf'))332 write_mol_with_coords(mol_pred, pos, os.path.join(write_dir, f'{orig_complex_graph.name[0]}_rmsd_{rmsd_i}_confidence_{score}.sdf'))333 ########################################################################334 335 centroid_distances_list.append(centroid_distance)336 337 cross_distances = np.linalg.norm(receptor_pos[:, :, None, :] - ligand_pos[:, None, :, :], axis=-1)338 min_cross_distances_list.append(np.min(cross_distances, axis=(1, 2)))339 self_distances = np.linalg.norm(ligand_pos[:, :, None, :] - ligand_pos[:, None, :, :], axis=-1)340 self_distances = np.where(np.eye(self_distances.shape[2]), np.inf, self_distances)341 min_self_distances_list.append(np.min(self_distances, axis=(1, 2)))342 343 base_cross_distances = np.linalg.norm(receptor_pos[:, :, None, :] - orig_ligand_pos[:, None, :, :], axis=-1)344 base_min_cross_distances_list.append(np.min(base_cross_distances, axis=(1, 2)))345 346 if args.save_visualisation:347 write_dir_vis = f'{args.out_dir}/docking_result_{os.path.basename(args.split_path)}/{orig_complex_graph.name[0]}'348 os.makedirs(write_dir, exist_ok=True)349 if confidence is not None:350 for rank, batch_idx in enumerate(re_order):351 try:352 visualization_list[batch_idx].write(353 f'{write_dir_vis}/{data_list[batch_idx]["name"][0]}_{rank + 1}_{rmsd[batch_idx]:.1f}_{(confidence)[batch_idx]:.1f}.pdb')354 except:355 continue356 else:357 for rank, batch_idx in enumerate(np.argsort(rmsd)):358 try:359 visualization_list[batch_idx].write(360 f'{write_dir_vis}/{data_list[batch_idx]["name"][0]}_{rank + 1}_{rmsd[batch_idx]:.1f}.pdb')361 except:362 continue363 without_rec_overlap_list.append(1 if orig_complex_graph.name[0] in names_no_rec_overlap else 0)364 names_list.append(name_map_idx[orig_complex_graph.name[0]])365 except Exception as e:366 logger.info("Failed on", orig_complex_graph["name"], e)367 failures += 1368 sample_count_failed +=1369 if sample_count_failed > 5:370 logger.info(" Skip by five times Failed on", orig_complex_graph["name"], e)371 success = 1372 else:373 success = 0374 375 accelerator.wait_for_everyone()376 rmsds_list, centroid_distances_list, failures, skipped, min_cross_distances_list, base_min_cross_distances_list, confidences_list =\377 accelerator.gather(torch.tensor(rmsds_list).to(device)), accelerator.gather(torch.tensor(centroid_distances_list).to(device)),accelerator.gather(torch.tensor(failures).to(device)),accelerator.gather(torch.tensor(skipped).to(device)),\378 accelerator.gather(torch.tensor(min_cross_distances_list).to(device)), accelerator.gather(torch.tensor(base_min_cross_distances_list).to(device)), accelerator.gather(torch.tensor(confidences_list).to(device))379 run_times, min_self_distances_list, without_rec_overlap_list = accelerator.gather(torch.tensor(run_times).to(device)), accelerator.gather(torch.tensor(min_self_distances_list).to(device)),accelerator.gather(torch.tensor(without_rec_overlap_list).to(device))380 rmsds_list, centroid_distances_list, failures, skipped, min_cross_distances_list, base_min_cross_distances_list, confidences_list = \381 rmsds_list.cpu().detach().numpy(), centroid_distances_list.cpu().detach().numpy(), failures.cpu().detach().numpy(), skipped.cpu().detach().numpy(), min_cross_distances_list.cpu().detach().numpy(), base_min_cross_distances_list.cpu().detach().numpy(), confidences_list.cpu().detach().numpy()382 run_times, min_self_distances_list, without_rec_overlap_list = \383 run_times.cpu().detach().numpy(), min_self_distances_list.cpu().detach().numpy(), without_rec_overlap_list.cpu().detach().numpy()384 names_list = accelerator.gather(torch.tensor(names_list).to(device))385 names_list = names_list.cpu().detach().numpy()386 accelerator.wait_for_everyone()387 if accelerator.is_local_main_process:388 389 logger.info('Performance without hydrogens included in the loss')390 logger.info(failures, "failures due to exceptions")391 logger.info(skipped, ' skipped because complex was not in confidence dataset')392 performance_metrics = {}393 for overlap in ['', 'no_overlap_']:394 if 'no_overlap_' == overlap:395 without_rec_overlap = np.array(without_rec_overlap_list, dtype=bool)396 if without_rec_overlap.sum() == 0: continue397 rmsds = np.array(rmsds_list)[without_rec_overlap]398 min_self_distances = np.array(min_self_distances_list)[without_rec_overlap]399 centroid_distances = np.array(centroid_distances_list)[without_rec_overlap]400 # if confidence_model is not None:401 confidences = np.array(confidences_list)[without_rec_overlap]402 # else:403 # confidences = None404 min_cross_distances = np.array(min_cross_distances_list)[without_rec_overlap]405 base_min_cross_distances = np.array(base_min_cross_distances_list)[without_rec_overlap]406 names = np.array(names_list)[without_rec_overlap]407 else:408 rmsds = np.array(rmsds_list)409 min_self_distances = np.array(min_self_distances_list)410 centroid_distances = np.array(centroid_distances_list)411 # if confidence_model is not None:412 confidences = np.array(confidences_list)413 # else:414 # confidences = None415 min_cross_distances = np.array(min_cross_distances_list)416 base_min_cross_distances = np.array(base_min_cross_distances_list)417 names = np.array(names_list)418 names = np.array([idx_map_name[idx] for idx in names])419 420 run_times = np.array(run_times)421 np.save(f'{args.out_dir}/{overlap}min_cross_distances.npy', min_cross_distances)422 np.save(f'{args.out_dir}/{overlap}min_self_distances.npy', min_self_distances)423 np.save(f'{args.out_dir}/{overlap}base_min_cross_distances.npy', base_min_cross_distances)424 np.save(f'{args.out_dir}/{overlap}rmsds.npy', rmsds)425 np.save(f'{args.out_dir}/{overlap}centroid_distances.npy', centroid_distances)426 np.save(f'{args.out_dir}/{overlap}confidences.npy', confidences)427 np.save(f'{args.out_dir}/{overlap}run_times.npy', run_times)428 np.save(f'{args.out_dir}/{overlap}complex_names.npy', np.array(names))429 430 performance_metrics.update({431 f'{overlap}run_times_std': run_times.std().__round__(2),432 f'{overlap}run_times_mean': run_times.mean().__round__(2),433 f'{overlap}steric_clash_fraction': (434 100 * (min_cross_distances < 0.4).sum() / len(min_cross_distances) / N).__round__(2),435 f'{overlap}self_intersect_fraction': (436 100 * (min_self_distances < 0.4).sum() / len(min_self_distances) / N).__round__(2),437 f'{overlap}mean_rmsd': rmsds.mean(),438 f'{overlap}rmsds_below_1': (100 * (rmsds < 1).sum() / len(rmsds) / N),439 f'{overlap}rmsds_below_2': (100 * (rmsds < 2).sum() / len(rmsds) / N),440 f'{overlap}rmsds_below_5': (100 * (rmsds < 5).sum() / len(rmsds) / N),441 f'{overlap}rmsds_percentile_25': np.percentile(rmsds, 25).round(2),442 f'{overlap}rmsds_percentile_50': np.percentile(rmsds, 50).round(2),443 f'{overlap}rmsds_percentile_75': np.percentile(rmsds, 75).round(2),444 445 f'{overlap}mean_centroid': centroid_distances.mean().__round__(2),446 f'{overlap}centroid_below_2': (100 * (centroid_distances < 2).sum() / len(centroid_distances) / N).__round__(2),447 f'{overlap}centroid_below_5': (100 * (centroid_distances < 5).sum() / len(centroid_distances) / N).__round__(2),448 f'{overlap}centroid_percentile_25': np.percentile(centroid_distances, 25).round(2),449 f'{overlap}centroid_percentile_50': np.percentile(centroid_distances, 50).round(2),450 f'{overlap}centroid_percentile_75': np.percentile(centroid_distances, 75).round(2),451 })452 453 if N >= 5:454 top5_rmsds = np.min(rmsds[:, :5], axis=1)455 top5_centroid_distances = centroid_distances[456 np.arange(rmsds.shape[0])[:, None], np.argsort(rmsds[:, :5], axis=1)][:, 0]457 top5_min_cross_distances = min_cross_distances[458 np.arange(rmsds.shape[0])[:, None], np.argsort(rmsds[:, :5], axis=1)][:, 0]459 top5_min_self_distances = min_self_distances[460 np.arange(rmsds.shape[0])[:, None], np.argsort(rmsds[:, :5], axis=1)][:, 0]461 performance_metrics.update({462 f'{overlap}top5_steric_clash_fraction': (463 100 * (top5_min_cross_distances < 0.4).sum() / len(top5_min_cross_distances)).__round__(2),464 f'{overlap}top5_self_intersect_fraction': (465 100 * (top5_min_self_distances < 0.4).sum() / len(top5_min_self_distances)).__round__(2),466 f'{overlap}top5_rmsds_below_1': (100 * (top5_rmsds < 1).sum() / len(top5_rmsds)).__round__(2),467 f'{overlap}top5_rmsds_below_2': (100 * (top5_rmsds < 2).sum() / len(top5_rmsds)).__round__(2),468 f'{overlap}top5_rmsds_below_5': (100 * (top5_rmsds < 5).sum() / len(top5_rmsds)).__round__(2),469 f'{overlap}top5_rmsds_percentile_25': np.percentile(top5_rmsds, 25).round(2),470 f'{overlap}top5_rmsds_percentile_50': np.percentile(top5_rmsds, 50).round(2),471 f'{overlap}top5_rmsds_percentile_75': np.percentile(top5_rmsds, 75).round(2),472 473 f'{overlap}top5_centroid_below_2': (474 100 * (top5_centroid_distances < 2).sum() / len(top5_centroid_distances)).__round__(2),475 f'{overlap}top5_centroid_below_5': (476 100 * (top5_centroid_distances < 5).sum() / len(top5_centroid_distances)).__round__(2),477 f'{overlap}top5_centroid_percentile_25': np.percentile(top5_centroid_distances, 25).round(2),478 f'{overlap}top5_centroid_percentile_50': np.percentile(top5_centroid_distances, 50).round(2),479 f'{overlap}top5_centroid_percentile_75': np.percentile(top5_centroid_distances, 75).round(2),480 })481 482 if N >= 10:483 top10_rmsds = np.min(rmsds[:, :10], axis=1)484 top10_centroid_distances = centroid_distances[485 np.arange(rmsds.shape[0])[:, None], np.argsort(rmsds[:, :10], axis=1)][:, 0]486 top10_min_cross_distances = min_cross_distances[487 np.arange(rmsds.shape[0])[:, None], np.argsort(rmsds[:, :10], axis=1)][:, 0]488 top10_min_self_distances = min_self_distances[489 np.arange(rmsds.shape[0])[:, None], np.argsort(rmsds[:, :10], axis=1)][:, 0]490 performance_metrics.update({491 f'{overlap}top10_steric_clash_fraction': (492 100 * (top10_min_cross_distances < 0.4).sum() / len(top10_min_cross_distances)).__round__(2),493 f'{overlap}top10_self_intersect_fraction': (494 100 * (top10_min_self_distances < 0.4).sum() / len(top10_min_self_distances)).__round__(2),495 f'{overlap}top10_rmsds_below_1': (100 * (top10_rmsds < 1).sum() / len(top10_rmsds)).__round__(2),496 f'{overlap}top10_rmsds_below_2': (100 * (top10_rmsds < 2).sum() / len(top10_rmsds)).__round__(2),497 f'{overlap}top10_rmsds_below_5': (100 * (top10_rmsds < 5).sum() / len(top10_rmsds)).__round__(2),498 f'{overlap}top10_rmsds_percentile_25': np.percentile(top10_rmsds, 25).round(2),499 f'{overlap}top10_rmsds_percentile_50': np.percentile(top10_rmsds, 50).round(2),500 f'{overlap}top10_rmsds_percentile_75': np.percentile(top10_rmsds, 75).round(2),501 502 f'{overlap}top10_centroid_below_2': (503 100 * (top10_centroid_distances < 2).sum() / len(top10_centroid_distances)).__round__(2),504 f'{overlap}top10_centroid_below_5': (505 100 * (top10_centroid_distances < 5).sum() / len(top10_centroid_distances)).__round__(2),506 f'{overlap}top10_centroid_percentile_25': np.percentile(top10_centroid_distances, 25).round(2),507 f'{overlap}top10_centroid_percentile_50': np.percentile(top10_centroid_distances, 50).round(2),508 f'{overlap}top10_centroid_percentile_75': np.percentile(top10_centroid_distances, 75).round(2),509 })510 511 # if confidence_model is not None:512 if confidences is not None:513 confidence_ordering = np.argsort(confidences, axis=1)[:, ::-1]514 515 filtered_rmsds = rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, 0]516 filtered_centroid_distances = centroid_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, 0]517 filtered_min_cross_distances = min_cross_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:,518 0]519 filtered_min_self_distances = min_self_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, 0]520 performance_metrics.update({521 f'{overlap}filtered_self_intersect_fraction': (522 100 * (filtered_min_self_distances < 0.4).sum() / len(filtered_min_self_distances)).__round__(523 2),524 f'{overlap}filtered_steric_clash_fraction': (525 100 * (filtered_min_cross_distances < 0.4).sum() / len(filtered_min_cross_distances)).__round__(526 2),527 f'{overlap}filtered_rmsds_below_1': (100 * (filtered_rmsds < 1).sum() / len(filtered_rmsds)).__round__(2),528 f'{overlap}filtered_rmsds_below_2': (100 * (filtered_rmsds < 2).sum() / len(filtered_rmsds)).__round__(2),529 f'{overlap}filtered_rmsds_below_5': (100 * (filtered_rmsds < 5).sum() / len(filtered_rmsds)).__round__(2),530 f'{overlap}filtered_rmsds_percentile_25': np.percentile(filtered_rmsds, 25).round(2),531 f'{overlap}filtered_rmsds_percentile_50': np.percentile(filtered_rmsds, 50).round(2),532 f'{overlap}filtered_rmsds_percentile_75': np.percentile(filtered_rmsds, 75).round(2),533 534 f'{overlap}filtered_centroid_below_2': (535 100 * (filtered_centroid_distances < 2).sum() / len(filtered_centroid_distances)).__round__(2),536 f'{overlap}filtered_centroid_below_5': (537 100 * (filtered_centroid_distances < 5).sum() / len(filtered_centroid_distances)).__round__(2),538 f'{overlap}filtered_centroid_percentile_25': np.percentile(filtered_centroid_distances, 25).round(2),539 f'{overlap}filtered_centroid_percentile_50': np.percentile(filtered_centroid_distances, 50).round(2),540 f'{overlap}filtered_centroid_percentile_75': np.percentile(filtered_centroid_distances, 75).round(2),541 })542 543 if N >= 5:544 top5_filtered_rmsds = np.min(rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5], axis=1)545 top5_filtered_centroid_distances = \546 centroid_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5][547 np.arange(rmsds.shape[0])[:, None], np.argsort(548 rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5], axis=1)][:, 0]549 top5_filtered_min_cross_distances = \550 min_cross_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5][551 np.arange(rmsds.shape[0])[:, None], np.argsort(552 rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5], axis=1)][:, 0]553 top5_filtered_min_self_distances = \554 min_self_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5][555 np.arange(rmsds.shape[0])[:, None], np.argsort(556 rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :5], axis=1)][:, 0]557 performance_metrics.update({558 f'{overlap}top5_filtered_self_intersect_fraction': (559 100 * (top5_filtered_min_cross_distances < 0.4).sum() / len(560 top5_filtered_min_cross_distances)).__round__(2),561 f'{overlap}top5_filtered_steric_clash_fraction': (562 100 * (top5_filtered_min_cross_distances < 0.4).sum() / len(563 top5_filtered_min_cross_distances)).__round__(2),564 f'{overlap}top5_filtered_rmsds_below_1': (565 100 * (top5_filtered_rmsds < 1).sum() / len(top5_filtered_rmsds)).__round__(2),566 f'{overlap}top5_filtered_rmsds_below_2': (567 100 * (top5_filtered_rmsds < 2).sum() / len(top5_filtered_rmsds)).__round__(2),568 f'{overlap}top5_filtered_rmsds_below_5': (569 100 * (top5_filtered_rmsds < 5).sum() / len(top5_filtered_rmsds)).__round__(2),570 f'{overlap}top5_filtered_rmsds_percentile_25': np.percentile(top5_filtered_rmsds, 25).round(2),571 f'{overlap}top5_filtered_rmsds_percentile_50': np.percentile(top5_filtered_rmsds, 50).round(2),572 f'{overlap}top5_filtered_rmsds_percentile_75': np.percentile(top5_filtered_rmsds, 75).round(2),573 574 f'{overlap}top5_filtered_centroid_below_2': (100 * (top5_filtered_centroid_distances < 2).sum() / len(575 top5_filtered_centroid_distances)).__round__(2),576 f'{overlap}top5_filtered_centroid_below_5': (100 * (top5_filtered_centroid_distances < 5).sum() / len(577 top5_filtered_centroid_distances)).__round__(2),578 f'{overlap}top5_filtered_centroid_percentile_25': np.percentile(top5_filtered_centroid_distances,579 25).round(2),580 f'{overlap}top5_filtered_centroid_percentile_50': np.percentile(top5_filtered_centroid_distances,581 50).round(2),582 f'{overlap}top5_filtered_centroid_percentile_75': np.percentile(top5_filtered_centroid_distances,583 75).round(2),584 })585 if N >= 10:586 top10_filtered_rmsds = np.min(rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10],587 axis=1)588 top10_filtered_centroid_distances = \589 centroid_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10][590 np.arange(rmsds.shape[0])[:, None], np.argsort(591 rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10], axis=1)][:, 0]592 top10_filtered_min_cross_distances = \593 min_cross_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10][594 np.arange(rmsds.shape[0])[:, None], np.argsort(595 rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10], axis=1)][:, 0]596 top10_filtered_min_self_distances = \597 min_self_distances[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10][598 np.arange(rmsds.shape[0])[:, None], np.argsort(599 rmsds[np.arange(rmsds.shape[0])[:, None], confidence_ordering][:, :10], axis=1)][:, 0]600 performance_metrics.update({601 f'{overlap}top10_filtered_self_intersect_fraction': (602 100 * (top10_filtered_min_cross_distances < 0.4).sum() / len(603 top10_filtered_min_cross_distances)).__round__(2),604 f'{overlap}top10_filtered_steric_clash_fraction': (605 100 * (top10_filtered_min_cross_distances < 0.4).sum() / len(606 top10_filtered_min_cross_distances)).__round__(2),607 f'{overlap}top10_filtered_rmsds_below_1': (608 100 * (top10_filtered_rmsds < 1).sum() / len(top10_filtered_rmsds)).__round__(2),609 f'{overlap}top10_filtered_rmsds_below_2': (610 100 * (top10_filtered_rmsds < 2).sum() / len(top10_filtered_rmsds)).__round__(2),611 f'{overlap}top10_filtered_rmsds_below_5': (612 100 * (top10_filtered_rmsds < 5).sum() / len(top10_filtered_rmsds)).__round__(2),613 f'{overlap}top10_filtered_rmsds_percentile_25': np.percentile(top10_filtered_rmsds, 25).round(2),614 f'{overlap}top10_filtered_rmsds_percentile_50': np.percentile(top10_filtered_rmsds, 50).round(2),615 f'{overlap}top10_filtered_rmsds_percentile_75': np.percentile(top10_filtered_rmsds, 75).round(2),616 617 f'{overlap}top10_filtered_centroid_below_2': (100 * (top10_filtered_centroid_distances < 2).sum() / len(618 top10_filtered_centroid_distances)).__round__(2),619 f'{overlap}top10_filtered_centroid_below_5': (100 * (top10_filtered_centroid_distances < 5).sum() / len(620 top10_filtered_centroid_distances)).__round__(2),621 f'{overlap}top10_filtered_centroid_percentile_25': np.percentile(top10_filtered_centroid_distances,622 25).round(2),623 f'{overlap}top10_filtered_centroid_percentile_50': np.percentile(top10_filtered_centroid_distances,624 50).round(2),625 f'{overlap}top10_filtered_centroid_percentile_75': np.percentile(top10_filtered_centroid_distances,626 75).round(2),627 })628 629 for k in performance_metrics:630 logger.info(k, performance_metrics[k])631 632 if args.wandb:633 wandb.log(performance_metrics)634 histogram_metrics_list = [('rmsd', rmsds[:, 0]),635 ('centroid_distance', centroid_distances[:, 0]),636 ('mean_rmsd', rmsds.mean(axis=1)),637 ('mean_centroid_distance', centroid_distances.mean(axis=1))]638 if N >= 5:639 histogram_metrics_list.append(('top5_rmsds', top5_rmsds))640 histogram_metrics_list.append(('top5_centroid_distances', top5_centroid_distances))641 if N >= 10:642 histogram_metrics_list.append(('top10_rmsds', top10_rmsds))643 histogram_metrics_list.append(('top10_centroid_distances', top10_centroid_distances))644 # if confidence_model is not None:645 if confidences is not None:646 histogram_metrics_list.append(('filtered_rmsd', filtered_rmsds))647 histogram_metrics_list.append(('filtered_centroid_distance', filtered_centroid_distances))648 if N >= 5:649 histogram_metrics_list.append(('top5_filtered_rmsds', top5_filtered_rmsds))650 histogram_metrics_list.append(('top5_filtered_centroid_distances', top5_filtered_centroid_distances))651 if N >= 10:652 histogram_metrics_list.append(('top10_filtered_rmsds', top10_filtered_rmsds))653 histogram_metrics_list.append(('top10_filtered_centroid_distances', top10_filtered_centroid_distances))654 if args.wandb:655 wandb.finish()656if __name__ == '__main__':657 from accelerate import Accelerator658 from accelerate.utils import DistributedDataParallelKwargs659 kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)660 accelerator = Accelerator(kwargs_handlers=[kwargs])661 from accelerate.utils import set_seed662 device = accelerator.device663 set_seed(1024)664 accelerator.logger.info(f'device {str(accelerator.device)} is used!')665 main_function()666 # sys.exit()667 