CoolFace
Modelpublic

OneScience-Group/SurfDock

sourceHugging Facemitupdated 17d agoView on Hugging Face
0likes25downloads
post_energy_minimize.py164 linesDownload Raw Back to force_optimize
1import os2from minimize_utils import GetfixedPDB,GetFFGenerator,UpdatePose,GetPlatformPara,GetPlatform,Molecule,trySystem,read_molecule,run_command,read_abs_file_mol3import sys4from openmm.app import Modeller5from joblib import Parallel,delayed6import argparse7from tqdm import tqdm8from glob import glob9import warnings10import traceback11import time12import pandas as pd13import numpy as np14import logging15 16 17""""18This Script will help user to do energy minimized for protein-ligand complex by openmm19Of course , you can use force_optimize args in docking step if you want to minimized all docking pose!,but may be it will be slowly20So , I think you can use this script to do energy minimized for protein-ligand complex that ranking topN in docking step,this will save more time,without performance loss21Enjoy it!22 23"""24if __name__ == '__main__':25    logging.basicConfig(level=logging.INFO)26    logger = logging.getLogger(__name__)27    parser = argparse.ArgumentParser(description='Process protein-ligand files.')28    parser.add_argument('--head_num', type=int, default=20, help='Number of top pose to be minimized.')29    parser.add_argument('--num_process', type=int, default=20, help='Number of parallel workers.')30    parser.add_argument('--cuda', type=int, default=0, help='Number of parallel workers.')31    parser.add_argument('--path_csv', type=str, default='~/Screen_dataset/dataset/DEKOIS2_SurfDock_pose.csv', help='path csv file')32    parser.add_argument('--out_dir', type=str, default='~/Screen_dataset/SurfDock_multi_pose_minimized', help='save_dir')33    parser.add_argument('--head_index', type=int, default=0, help='the head index to start minimized,this optinal to minimized use multi-GPU every GPU minimized a part of sdfs')34    parser.add_argument('--tail_index', type=int, default=-1, help='the tail index to start minimized,this optinal to minimized use multi-GPU every GPU minimized a part of sdfs')35    args = parser.parse_args()36    os.environ['OMP_NUM_THREADS'] = '1'37    """Init force field"""38    start_time = time.time()39    platform = GetPlatformPara()40    system_generator = GetFFGenerator(ignoreExternalBonds=True)41    system_generator_gaff = GetFFGenerator(small_molecule_forcefield = 'gaff-2.11',ignoreExternalBonds=True)42    paths = pd.read_csv(args.path_csv)43    for protein_path,sdf_dir in zip(paths['protein_path'],paths['ligand_path']):44        pdbid = os.path.basename(protein_path).split('_')[0]45        try:46            logger.info(f'minimized for target {pdbid}.......')47            logger.info('Use default forcefield')48            receptor_path = protein_path49            fixer = GetfixedPDB(receptor_path)50            modeller = Modeller(fixer.topology, fixer.positions)51            if os.path.isdir(sdf_dir):52                logger.info(f" {sdf_dir} is a Dir path,if you want to minimized just a file like relax for esmfold-ligand complex,please check the ligand_path !")53                # pass54                sdf_paths = glob(os.path.join(sdf_dir, '*.sdf'))55                # this code use to energy minimized the top N pose for every molecule56                # select confidence topN pose to minimized57                sdf_pd = pd.DataFrame({'pred_sdf_name':sdf_paths})58                sdf_pd['molecule_name'] = sdf_pd['pred_sdf_name'].apply(lambda x: os.path.basename(x).split('_sample_idx_')[0])59                sdf_pd['confidence'] = sdf_pd['pred_sdf_name'].apply(lambda x: float(os.path.basename(x).split('_confidence_')[-1].split('.sdf')[0]))60                # selected the topN confidence pose61                result = sdf_pd.sort_values('confidence',ascending=False)62                result_group = result.groupby('molecule_name')63                result = result_group.head(args.head_num)64                top1_sdfs = result['pred_sdf_name'].tolist()[args.head_index:args.tail_index]65            else:66                logger.info(f"Only minimized file {sdf_dir},if you want to minimized docking result from a Dir ,please check the ligand_path !")67                logger.info(f"Only minimized file {sdf_dir},head_num,head_index, tail_index, out_dir params will unable!")68                args.out_dir = os.path.dirname(os.path.dirname(sdf_dir))69                top1_sdfs = [sdf_dir]70            71            72            logger.info(f"ALL About {len(top1_sdfs)} sdfs to minimize , try to skip files have done!")73            if os.path.isdir(sdf_dir):74                # check out_dir done have optimized file and filter optimized files75                if os.path.exists(os.path.join(args.out_dir ,os.path.basename(sdf_dir))):76                    finished_files=os.listdir(os.path.join(args.out_dir ,os.path.basename(sdf_dir)))77                else:78                    finished_files = []79                if os.path.exists(os.path.join(args.out_dir ,os.path.basename(sdf_dir) + '_tmp')):80                    finished_files.extend(os.listdir(os.path.join(args.out_dir ,os.path.basename(sdf_dir) + '_tmp')))81 82                top1_sdfs = list(filter(lambda x:os.path.splitext(os.path.basename(x))[0]+ '_minimized.sdf' not in finished_files and \83                    os.path.splitext(os.path.basename(x))[0]+ '_unminimized.sdf' not in finished_files84                    ,top1_sdfs))85            else:86                87                if os.path.exists(os.path.splitext(os.path.basename(top1_sdfs[0]))[0] + '_minimized.sdf') or os.path.exists(os.path.splitext(os.path.basename(top1_sdfs[0]))[0] + '_unminimized.sdf') :88                    logger.info(f"{os.path.splitext(os.path.basename(top1_sdfs[0]))[0]} have been minimized,skip it!")89                    # finished_files = []90                    continue91                else:92                    finished_files = []93 94               95 96            logger.info(f"Minimizeing...... {len(finished_files)} sdfs have Minimized || left {len(top1_sdfs)} sdfs to Minimizing......")97 98            logger.info(f"Trying...... create system for protein!")99            failed_create_system = False100            101            for test_idx in range(len(top1_sdfs)):102                try:103                    dockingpose = read_abs_file_mol(top1_sdfs[test_idx], remove_hs=True, sanitize=True)104                    lig_mol = Molecule.from_rdkit(dockingpose,allow_undefined_stereo=True)105                    # set formal_charge use gasteiger106                    lig_mol.assign_partial_charges(partial_charge_method='gasteiger')107                    modeller = trySystem(system_generator_gaff,modeller,lig_mol,top1_sdfs[test_idx])108                    failed_create_system = False109                    break110                except:111                    logger.info(f"ERROR in create system step! try anather molecule ing....., or you can check the protein please!")112                    failed_create_system = True113                    continue114            115            if failed_create_system:116                logger.info(f"ERROR For create system for protein!,check in error_for_create_system.txt")117                with open('error_for_create_system.txt','a') as f:118                    f.write(receptor_path +': Create system error! by :' + '\n')119                continue120 121 122            if modeller is None:123                print('Create system error!')124                with open('error_for_create_system.txt','a') as f:125                    f.write(receptor_path +': Create system error! by :' + '\n')126                logger.info(f"ERROR For create system for protein!,check in error_for_create_system.txt")127                continue128            logger.info(f"Done For create system for protein!,Start to Minimize sdf file")129            130            protein_atoms = list(modeller.topology.atoms())131 132            with Parallel(n_jobs=args.num_process,) as parallel:133                new_data_list = parallel(delayed(UpdatePose)(lig_path,system_generator,modeller,protein_atoms,args.out_dir) for lig_path in top1_sdfs)134            # selected the failed samples and try to use gaff-2.11 forcefield135            if sum(new_data_list) != 0:136                result = np.array(new_data_list)137                indices = np.where(result == 1)138                failed_sdfs = [top1_sdfs[i] for i in indices[0]]139                logger.info(f'Minimized not Completed:{pdbid}, {len(failed_sdfs)} sdf not be minimized by default forcefield , try use gaff-2.11 forcefield!')140                with Parallel(n_jobs=args.num_process) as parallel:141                    new_data_list = parallel(delayed(UpdatePose)(lig_path,system_generator_gaff,modeller,protein_atoms,args.out_dir) for lig_path in failed_sdfs)142            143            if sum(new_data_list) != 0:144                145                logger.info(f'Minimized not Completed:{pdbid}, {sum(new_data_list)} sdf not be minimized,use unminimized conformers for later stage')146                # save unminimized conformers147                result = np.array(new_data_list)148                indices = np.where(result == 1)149                failed_sdfs = [top1_sdfs[i] for i in indices[0]]150                out_base_dir = os.path.join(args.out_dir,failed_sdfs[0].split('/')[-2])151                cwd_path = os.path.dirname(os.path.abspath(__file__))152                os.makedirs(out_base_dir,exist_ok=True)153                for lig_path in failed_sdfs:154                    out_file = os.path.join(out_base_dir,os.path.splitext(os.path.basename(lig_path))[0] + '_unminimized.sdf')155                    command = f"cp {lig_path} {out_file}"156                    run_command(command=command,cwd_path = cwd_path )157            logger.info(f'Finish minimized target {pdbid}')158        except Exception as e:159            warnings.warn(f'{pdbid} faild with {str(e)}')160            error_info = traceback.format_exc()161            print(error_info)162    end_time = time.time()163    logger.info(f"Time taken for optimizing {len(paths)} molecules: {end_time - start_time:.2f} seconds")164