CoolFace
Modelpublic

OneScience-Group/SurfDock

sourceHugging Facemitupdated 15d agoView on Hugging Face
0likes25downloads
inference_accelerate.py480 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 gc22import pandas as pd23import wandb24import glob25from rdkit import RDLogger26from rdkit.Chem import RemoveHs27from datasets.process_mols import write_mol_with_coords28from torch_geometric.loader import DataLoader29from datasets.pdbbind import PDBBind, read_mol,read_abs_file_mol30from utils.diffusion_utils import t_to_sigma as t_to_sigma_compl, get_t_schedule31from utils.sampling import randomize_position, sampling,inferenceFFOptimize32from utils.utils import get_symmetry_rmsd, remove_all_hs33from score_in_place_dataset.score_dataset import ScreenDataset34from utils.utils import get_model, ExponentialMovingAverage35from utils.visualise import PDBFile36from tqdm import tqdm37from collections import defaultdict38from packaging import version39import warnings40warnings.filterwarnings("ignore", category=UserWarning, module="torch.jit")41RDLogger.DisableLog('rdApp.*')42import yaml43from loguru import logger44 45cache_name = datetime.now().strftime('date%d-%m_time%H-%M-%S.%f')46parser = ArgumentParser()47 48parser.add_argument('--config', type=FileType(mode='r'), default=None)49parser.add_argument('--data_csv', type=str, default='~/Screen_dataset/dataset/DEKOIS2.csv', help='Path to folder with dataset for score in place')50parser.add_argument('--model_dir', type=str, default=None, help='Path to folder with trained score model and hyperparameters')51parser.add_argument('--ckpt', type=str, default=None, help='Checkpoint to use inside the folder')52parser.add_argument('--confidence_model_dir', type=str, default=None, help='Path to folder with trained confidence model and hyperparameters')53parser.add_argument('--confidence_ckpt', type=str, default=None, help='Checkpoint to use inside the folder')54# save docking result or not55parser.add_argument('--save_docking_result', action='store_true', default=False, help='Whether to save docking result')56# put ligand to pocket center57parser.add_argument('--ligand_to_pocket_center', action='store_true', default=False, help='Whether to put ligand on pocket center')58parser.add_argument('--keep_input_pose', action='store_false', default=False, help='Whether keep original input pose')59parser.add_argument('--use_noise_to_rank', action='store_true', default=False, help='Whether to run the probability flow ODE')60parser.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.')61parser.add_argument('--run_name', type=str, default='test_ns_48_nv_10_layer_62023-06-25_07-54-08_model', help='')62parser.add_argument('--project', type=str, default='ligbind_inf_test_mdn', help='')63parser.add_argument('--surface_path', type=str, default='~/PDBBind_processed_8A_surface/', help='test dataset surface path')64parser.add_argument('--esm_embeddings_path', type=str, default='~/PDBBIND/esm_embedding/esm_embedding_pocket_for_train/esm2_3billion_embeddings.pt', help='test dataset esmbedding path')65parser.add_argument('--out_dir', type=str, default='~/test_workdir/mdn_result_40', help='Where to save results to')66parser.add_argument('--batch_size', type=int, default=40, help='Number of poses to sample in parallel we recommand set number = batch_size_molecule*samples_per_complex')67parser.add_argument('--batch_size_molecule', type=int, default=1, help='Number of molecul to sample in parallel')68parser.add_argument('--cache_path', type=str, default='~/PDBBIND/cache_PDBBIND_pocket_8A', help='Folder from where to load/restore cached dataset')69parser.add_argument('--data_dir', type=str, default='~/PDBBIND/PDBBind_pocket_8A/', help='Folder containing original structures')70parser.add_argument('--split_path', type=str, default='~/data/splits/timesplit_test', help='Path of file defining the split')71parser.add_argument('--no_overlap_names_path', type=str, default='~/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')72parser.add_argument('--no_model', action='store_true', default=False, help='Whether to return seed conformer without running model')73parser.add_argument('--no_random', action='store_true', default=False, help='Whether to add randomness in diffusion steps')74parser.add_argument('--no_final_step_noise', action='store_true', default=False, help='Whether to add noise after the final step')75parser.add_argument('--ode', action='store_true', default=False, help='Whether to run the probability flow ODE')76parser.add_argument('--wandb', action='store_true', default=False, help='')77parser.add_argument('--wandb_dir', type=str, default='~/test_workdir', help='Folder in which to save wandb logs')78parser.add_argument('--inference_steps', type=int, default=20, help='Number of denoising steps')79parser.add_argument('--limit_complexes', type=int, default=0, help='Limit to the number of complexes')80parser.add_argument('--num_workers', type=int, default=1, help='Number of workers for dataset creation')81parser.add_argument('--num_process', type=int, default=20, help='Number of parallel workers for minimized.')82parser.add_argument('--tqdm', action='store_true', default=False, help='Whether to show progress bar')83parser.add_argument('--save_visualisation', action='store_true', default=False, help='Whether to save visualizations')84parser.add_argument('--samples_per_complex', type=int, default=40, help='Number of poses to sample for each complex')85parser.add_argument('--save_docking_result_number', type=int, default=1, help='Number of poses to save in disk for each complex')86parser.add_argument('--actual_steps', type=int, default=None, help='')87parser.add_argument('--inference_mode', default='Screen', help='inference mode',choices=['Screen','evaluate'])88parser.add_argument('--head_index', type=int, default=0, help='the head index to start inference,this optinal to inference use multi-GPU every GPU minimized a part of csv file ')89parser.add_argument('--tail_index', type=int, default=-1, help='the tail index to start inference,this optinal to inference use multi-GPU every GPU minimized a part of csv file')90parser.add_argument('--ligandsMaxAtoms', type=int, default=80, help='the max number of atoms in ligand')91parser.add_argument('--random_seed', type=int, default=42,  help='random seed')92# force_minimized param93parser.add_argument('--force_optimize', action='store_true', default=False, help='')94parser.add_argument('--mdn_dist_threshold_test', type=float, default=3.0, help='mdn_dist_threshold_test')95args = parser.parse_args()96nowtime = datetime.now().strftime('%Y-%m-%d')97log_file_flag = '-'.join(args.project.split('/'))98logger.add(f'{os.path.dirname(args.out_dir)}/log-inference-{log_file_flag}-{nowtime}.log', rotation="500MB")99logger.info('Runing inference script in path: {}',os.getcwd())100logger.info('Runing inference with args: {}',args)101 102def main_function():103    if accelerator.is_local_main_process:104        if args.wandb:105            wandb.login(key = 'yourkey')106            run = wandb.init(107                entity='SurfDock',108                settings=wandb.Settings(start_method="fork"),109                project=args.project,110                name=args.run_name,111                dir = args.wandb_dir,112                config=args113            )114    if args.config:115        config_dict = yaml.load(args.config, Loader=yaml.FullLoader)116        arg_dict = args.__dict__117        for key, value in config_dict.items():118            if isinstance(value, list):119                for v in value:120                    arg_dict[key].append(v)121            else:122                arg_dict[key] = value123    if args.out_dir is None: args.out_dir = f'inference_out_dir_not_specified/{args.run_name}'124    os.makedirs(args.out_dir, exist_ok=True)125    with open(f'{args.model_dir}/model_parameters.yml') as f:126        score_model_args = Namespace(**yaml.full_load(f))127       128 129    if args.confidence_model_dir is not None:130        with open(f'{args.confidence_model_dir}/model_parameters.yml') as f:131            confidence_args = Namespace(**yaml.full_load(f))132            # 133            confidence_args.transfer_weights = False134            confidence_args.use_original_model_cache = True135            confidence_args.original_model_dir = None136            confidence_args.mdn_dist_threshold_test = args.mdn_dist_threshold_test if args.mdn_dist_threshold_test is not None else 5.0137            if not hasattr(confidence_args,'mdn_dist_threshold_train'):138                confidence_args.mdn_dist_threshold_train =7.0139 140    if args.confidence_model_dir is not None:141        if not (confidence_args.use_original_model_cache or confidence_args.transfer_weights):142            # 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 complexes143            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.')144            confidence_test_dataset = PDBBind(transform=None, root=args.data_dir, limit_complexes=args.limit_complexes,145                                    receptor_radius=confidence_args.receptor_radius,146                                cache_path=args.cache_path, split_path=args.split_path,147                                remove_hs=confidence_args.remove_hs, max_lig_size=None, c_alpha_max_neighbors=confidence_args.c_alpha_max_neighbors,148                                matching=not confidence_args.no_torsion, keep_original=True,149                                popsize=confidence_args.matching_popsize,150                                maxiter=confidence_args.matching_maxiter,151                                all_atoms=confidence_args.all_atoms,152                                atom_radius=confidence_args.atom_radius,153                                atom_max_neighbors=confidence_args.atom_max_neighbors,154                                esm_embeddings_path= args.esm_embeddings_path, require_ligand=True,155                                num_workers=args.num_workers,surface_path = args.surface_path)156            confidence_complex_dict = {d.name: d for d in confidence_test_dataset}157 158    t_to_sigma = partial(t_to_sigma_compl, args=score_model_args)159 160    if not args.no_model:161        model = get_model(score_model_args, device, t_to_sigma=t_to_sigma, no_parallel=True,model_type = score_model_args.model_type)162        state_dict = torch.load(f'{args.model_dir}/{args.ckpt}', map_location=torch.device('cpu'))163        if args.ckpt == 'last_model.pt':164            model_state_dict = state_dict['model']165            ema_weights_state = state_dict['ema_weights']166            model.load_state_dict(model_state_dict, strict=True)167            ema_weights = ExponentialMovingAverage(model.parameters(), decay=score_model_args.ema_rate)168            ema_weights.load_state_dict(ema_weights_state, device=device)169            ema_weights.copy_to(model.parameters())170        else:171            model.load_state_dict(state_dict, strict=False)172            model = model.to(device)173            model.eval()174            logger.info('loaded model weight for score model')175        if args.confidence_model_dir is not None:176            if confidence_args.transfer_weights:177                with open(f'{confidence_args.original_model_dir}/model_parameters.yml') as f:178                    confidence_model_args = Namespace(**yaml.full_load(f))179            else:180                confidence_model_args = confidence_args181 182            confidence_model = get_model(confidence_model_args, device, t_to_sigma=t_to_sigma, no_parallel=True,183                                        model_type = confidence_model_args.model_type)184            state_dict = torch.load(f'{args.confidence_model_dir}/{args.confidence_ckpt}', map_location=torch.device('cpu'))185            confidence_model.load_state_dict(state_dict, strict=True)186            confidence_model = confidence_model.to(device)187            confidence_model.eval()188        else:189            confidence_model = None190            confidence_args = None191            confidence_model_args = None192 193 194    tr_schedule = get_t_schedule(inference_steps=args.inference_steps)195    rot_schedule = tr_schedule196    tor_schedule = tr_schedule197    logger.info('t schedule:{}',tr_schedule)198    logger.info('Loading data ...........')199 200    """201    Load data from csv file to get the path of pocket,ligand,ref_ligand,surface202    203    """204    df = pd.read_csv(args.data_csv)[args.head_index:args.tail_index]205    protein_paths = df['protein_path'].tolist()206    pocket_paths = df['pocket_path'].tolist()207    ligands_paths = df['ligand_path'].tolist()208    ref_ligands = df['ref_ligand'].tolist()209    surface_paths = df['protein_surface'].tolist()210    if 'pocket_center' in df.columns:211        pocket_centers = df['pocket_center'].tolist()212        new_pocket_centers = []213        for center in pocket_centers:214            x = center.split(',')[0]215            y = center.split(',')[1]216            z = center.split(',')[2]217            new_pocket_centers.append(np.array([(float(x),float(y),float(z))]))218        pocket_centers = new_pocket_centers219    else:220        pocket_centers = [None]*len(protein_paths)221 222    esm_embeddings_dict = torch.load(args.esm_embeddings_path)223    confidence_list = []224    confidence_names = []225    sdf_names = []226    pocket_path_list =[]227 228    failures = 0229    N = args.samples_per_complex230    all_molecules = 0231    pbar = tqdm(zip(pocket_paths,ligands_paths,ref_ligands,surface_paths,protein_paths,pocket_centers),total=len(pocket_paths))232    start_time = time.time()233    for pocket_path,ligands_path,ref_ligand,surface_path,protein_path,pocket_center in pbar:234        in_loop_start_time = time.time()235        236        try:237        238            dirname = os.path.splitext(pocket_path.split('/')[-1])[0] + '_'+ os.path.splitext(ligands_path.split('/')[-1])[0] 239            write_dir =  os.path.join(args.out_dir,'SurfDock_docking_result',dirname)#f'{args.out_dir}/SurfDock_docking_result/{dirname}'240            os.makedirs(write_dir, exist_ok=True)241            242            esm_embeddings = copy.deepcopy(esm_embeddings_dict[os.path.splitext(os.path.basename(pocket_path))[0]])243 244            test_dataset = ScreenDataset(pocket_path,ligands_path,ref_ligand,surface_path,pocket_center,transform=None,245                                receptor_radius=confidence_args.receptor_radius,246                                cache_path=None, split_path=None,247                                remove_hs=confidence_args.remove_hs, max_lig_size=None,248                                c_alpha_max_neighbors=confidence_args.c_alpha_max_neighbors,249                                matching= False, keep_original=True,250                                popsize=confidence_args.matching_popsize,251                                maxiter=confidence_args.matching_maxiter,252                                all_atoms=confidence_args.all_atoms,253                                atom_radius=confidence_args.atom_radius,254                                atom_max_neighbors=confidence_args.atom_max_neighbors,255                                esm_embeddings=esm_embeddings,256                                require_ligand=False,257                                num_workers=args.num_workers,258                                keep_input_pose = args.keep_input_pose,259                                save_dir = write_dir,260                                inference_mode = args.inference_mode,261                                ligandsMaxAtoms=args.ligandsMaxAtoms)262            test_sample_num = len(test_dataset)263            all_molecules += test_sample_num264            test_loader = DataLoader(dataset=test_dataset, batch_size=args.batch_size_molecule, shuffle=False)265            if test_sample_num == 0:266                logger.error('No complexes need to be docking (skip before done or some errors) in {}', pocket_path)267                continue268            # test_loader= accelerator.prepare(test_loader)269            logger.info('Protein {} Size of test dataset: {}',os.path.splitext(os.path.basename(pocket_path))[0],  test_sample_num)270            ##### use torch.__version__orch complie to speed up the process ###271            # if version.parse(torch.__version__.split('+')[0])> version.parse("2.0"):272            #     model = torch.compile(model)273            #     confidence_model = torch.compile(confidence_model)274            #     logger.info('Your are using torch version={} , so SurfDock will use torch.compile to complie  model and confidence model',torch.__version__)275            #########################################################276            model = accelerator.prepare(model)277            test_loader= accelerator.prepare(test_loader)278            confidence_model = accelerator.prepare(confidence_model)279            """280            Start sampling conformers by SurfDock281            """282            for idx, orig_complex_graph in tqdm(enumerate(test_loader),total = len(test_loader),disable= not accelerator.is_local_main_process):283                284                try:285                    if 'ligand' not in orig_complex_graph.node_types:286                        logger.error('some error failed for conformer generate in rdkit: idx in batch graph: {}, ligand_path: {}',idx,ligands_path)287                        continue288                    orig_complex_graph_list = orig_complex_graph.to_data_list()289                    # add protein pocket information for minimized stage 290                    for temp_graph in orig_complex_graph_list:291                        temp_graph['protein_path'] = protein_path292                        temp_graph['pocket_path'] = pocket_path293 294                    success = 0295                    sample_count_failed = 0296                    data_list = []297                    # object 298                    data_list = [copy.deepcopy(temp_graph) for temp_graph in orig_complex_graph_list for _ in range(N)]299                    while not success: # keep trying in case of failure (sometimes stochastic)300                        301                        try:302                           303                            # data_list = [copy.deepcopy(temp_graph) for temp_graph in orig_complex_graph_list for _ in range(N)]304                            success = 1305                            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)306                            pdb = None307                            if args.save_visualisation:308                                visualization_list = []309                                for idx, graph in enumerate(data_list):310                                    # raw pose311                                    lig = read_mol(args.data_dir, graph['name'][0], remove_hs=score_model_args.remove_hs)312                                    pdb = PDBFile(lig)313                                    pdb.add(lig, 0, 0)314                                    # pose rdkit matching315                                    orig_complex_count = idx//N316                                         317                                    pdb.add((orig_complex_graph_list[orig_complex_count]['ligand'].pos + orig_complex_graph_list[orig_complex_count].original_center).detach().cpu(), 1, 0)318                                    # random rdkit matching319                                    pdb.add((graph['ligand'].pos + (graph.original_center).detach().cpu()), part=1, order=1)320                                    visualization_list.append(pdb)321                            else:322                                visualization_list = None323 324                            if not args.no_model:325 326                                confidence_data_list = None327 328                                data_list, confidence = sampling(input_data_list=data_list, model=model,329                                                                inference_steps=args.actual_steps if args.actual_steps is not None else args.inference_steps,330                                                                tr_schedule=tr_schedule, rot_schedule=rot_schedule,331                                                                tor_schedule=tor_schedule,332                                                                device=device, t_to_sigma=t_to_sigma, model_args=score_model_args,333                                                                no_random=args.no_random,334                                                                ode=args.ode, visualization_list=visualization_list,335                                                                confidence_model=confidence_model,336                                                                confidence_data_list=confidence_data_list,337                                                                confidence_model_args=confidence_model_args,338                                                                batch_size=args.batch_size,339                                                                no_final_step_noise=args.no_final_step_noise,args = args)340                                accelerator.wait_for_everyone()341 342                                confidence = confidence.cpu().detach().numpy()343 344                                # save confidence345                                confidence_list += confidence.tolist()346                                for _ in range(len(orig_complex_graph_list)):347                                    348                                    confidence_names.extend([orig_complex_graph_list[_]['name']]*N)349                                    pocket_path_list.extend([os.path.basename(pocket_path)]*N)350                                351                                sdf_names += [os.path.basename(ligands_path)]*len(confidence)352 353                                assert len(confidence_list)==len(confidence_names)==len(sdf_names)==len(pocket_path_list)354                            """ add a save command by caoduanhua to save the last state of ligand """355                            ########################################################################356                            if args.save_docking_result:357                                """"if you use multiple molecule parallel inference, you should re_order the confidence one by one"""358                                # add a parm to control the number of save ligand pose359                                head_threshold = 0360                                tail_threshold = N361                                confidence_tmp = confidence[head_threshold:tail_threshold]362                                re_order = np.argsort(confidence_tmp)[::-1]363                                if args.inference_mode=='evaluate':364                                    true_mol = remove_all_hs(read_abs_file_mol(ref_ligand))365                                for _ in range(len(orig_complex_graph_list)):366                                    for rank, batch_idx in enumerate(re_order[:args.save_docking_result_number]):367                                        true_idx = head_threshold + batch_idx368                                        mol_pred = copy.deepcopy(data_list[true_idx]['mol'])369                                        370                                        pos = data_list[true_idx]['ligand'].pos.cpu().numpy() + orig_complex_graph_list[_].original_center.cpu().numpy()371 372                                        if score_model_args.remove_hs: mol_pred = remove_all_hs(mol_pred)373                                        374                                        375                                        if args.inference_mode=='evaluate':376                                            try:377                                                rmsd = get_symmetry_rmsd(true_mol, true_mol.GetConformers()[0].GetPositions(), [pos])[0]378                                            except Exception as e:379                                                logger.warning("Using non corrected RMSD because of the error:{}", e)380 381                                                rmsd = np.sqrt(((true_mol.GetConformers()[0].GetPositions() - pos) ** 2).sum(axis=-1).mean(axis=0))382                                            result_filename = f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_rmsd_{rmsd}_confidence_{confidence_tmp[batch_idx]}.sdf'383                                        else:384                                            result_filename = f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_confidence_{confidence_tmp[batch_idx]}.sdf'385                                        386                                        write_mol_with_coords(mol_pred, pos, os.path.join(write_dir, result_filename))387                                        388                                        if args.save_visualisation:389                                            write_dir_vis = f'{args.out_dir}/SurfDock_docking_result/{data_list[true_idx]["name"]}'390                                            os.makedirs(write_dir, exist_ok=True)391                                            if args.inference_mode=='evaluate':392                                                vis_filename =f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_rmsd_{rmsd}_confidence_{confidence_tmp[batch_idx]}.pdb'393                                            else:394                                                vis_filename = f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_confidence_{confidence_tmp[batch_idx]}.pdb'395                                            try:396                                                visualization_list[batch_idx].write(397                                                    f'{write_dir_vis}/{vis_filename}')398                                            except:399                                                continue400                                    head_threshold += N401                                    tail_threshold += N402 403                                    if _ < len(orig_complex_graph_list) - 1:404  405                                        confidence_tmp = confidence[head_threshold:tail_threshold]406                                        re_order = np.argsort(confidence_tmp)[::-1]407                        except Exception as e:408                            # if isinstance(e,RecursionError) or 'out of memory' in str(e):409                            data_list = None410                            referrers = gc.get_referrers(data_list)411                            for ref in referrers:412                                    ref=None413                            gc.collect()414                            torch.cuda.empty_cache()415                            data_list = [copy.deepcopy(temp_graph) for temp_graph in orig_complex_graph_list for _ in range(N)]416                            logger.error("Failed on :{}, error of :{}", orig_complex_graph["name"], e)417                            failures += 1418                            sample_count_failed +=1419                            if sample_count_failed > 5:420                                logger.error(" Skip by five times Failed on :{}, error of :{}", orig_complex_graph["name"], e)421                                success = 1422                            else:423                                success = 0424 425                except Exception as e:426                    if 'out of memory' in str(e):427                        logger.critical('| WARNING: ran out of memory, skipping batch')428                    orig_complex_graph_list,orig_complex_graph,data_list=None,None,None429                    referrers = gc.get_referrers(data_list)430                    for ref in referrers:431                        ref=None432                    gc.collect()433                    torch.cuda.empty_cache()434                    logger.error('Some error failed for sampling: idx in batch : {}, ligand_path: {},error of :{} ',idx,ligands_path,e)435                    436                    continue437            # if args.inference_mode=='evaluate':438            esm_embeddings,test_dataset,test_loader,orig_complex_graph_list,orig_complex_graph,data_list=None,None,None,None,None,None439            gc.collect()440            torch.cuda.empty_cache()441        except Exception as e:442            logger.error('Some error failed for graph data. ligand_path: {},error of :{}',ligands_path,e)443            esm_embeddings,test_dataset,test_loader,orig_complex_graph_list,orig_complex_graph,data_list=None,None,None,None,None,None444            referrers = gc.get_referrers(data_list)445            for ref in referrers:446                ref=None447            gc.collect()448            torch.cuda.empty_cache()449            continue450        logger.info('Protein {} used time: {}',os.path.splitext(os.path.basename(pocket_path))[0],time.time() - in_loop_start_time)451    accelerator.wait_for_everyone()452    docking_time = time.time() - start_time453    if accelerator.is_local_main_process:454        logger.info('Docking time used for one moleculer: {}',docking_time/ all_molecules)455        logger.info('Docking time used: {}', docking_time)456        logger.info('Sampling conformers number: {}',all_molecules*args.samples_per_complex)457        logger.info('Output conformers number: {}',  all_molecules*args.save_docking_result_number)458        logger.info('Docking output molecule number: {}',all_molecules)459        # logger.info('Docking time used for one moleculer: {}',docking_time/all_molecules)460 461    result = pd.DataFrame({'sdf_name':sdf_names,'confidence':confidence_list,'confidence_name':confidence_names,'pocket_path':pocket_path_list})462    csv_flag = os.path.basename(args.data_csv).split('.')[0]463    result.to_csv(f'{args.out_dir}/{csv_flag}_head_{str(args.head_index)}_tail_{str(args.tail_index)}_confidence_on_device_{device}.csv',index=False)464 465    if accelerator.is_local_main_process:466        if args.wandb:467            wandb.finish()468if __name__ == '__main__':469    from accelerate import Accelerator470    from accelerate.utils import DistributedDataParallelKwargs471    kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)472    accelerator = Accelerator(kwargs_handlers=[kwargs])473    from accelerate.utils import set_seed474    device = accelerator.device475    set_seed(args.random_seed)476    from functools import partial477 478    accelerator.print(f'device {str(accelerator.device)} is used!')479    main_function()480