CoolFace
Modelpublic

OneScience-Group/SurfDock

sourceHugging Facemitupdated 15d agoView on Hugging Face
0likes25downloads
evaluate_accelarate.py667 linesDownload Raw Back to scripts
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