Download scripts/evaluate_score_in_place.py from OneScience-Group/SurfDock: direct link, hf CLI and curl.
- Browser
- Download file 13.6 kB
-
https://huggingface.co/OneScience-Group/SurfDock/resolve/main/scripts/evaluate_score_in_place.py
- Command line
-
hf download hf://OneScience-Group/SurfDock/scripts/evaluate_score_in_place.py
-
curl -L -o evaluate_score_in_place.py https://huggingface.co/OneScience-Group/SurfDock/resolve/main/scripts/evaluate_score_in_place.py
13.6 kB
| """ | |
| caoduanhua : we should to implemented a parapllel version of evaluate.py for a large dataset | |
| """ | |
| import copy | |
| import os | |
| import sys | |
| SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) | |
| PROJECT_DIR = os.path.dirname(SCRIPT_DIR) | |
| MODEL_DIR = os.path.join(PROJECT_DIR, "model") | |
| if MODEL_DIR not in sys.path: | |
| sys.path.insert(0, MODEL_DIR) | |
| import torch | |
| from argparse import ArgumentParser, Namespace, FileType | |
| from datetime import datetime | |
| import time | |
| import numpy as np | |
| import pandas as pd | |
| import wandb | |
| from rdkit import RDLogger | |
| from torch_geometric.loader import DataLoader | |
| from score_in_place_dataset.score_dataset import ScreenDataset | |
| # from utils.sampling import randomize_position, sampling | |
| from utils.utils import get_model | |
| from tqdm import tqdm | |
| from loguru import logger | |
| import warnings | |
| warnings.filterwarnings("ignore", category=UserWarning, module="torch.jit") | |
| RDLogger.DisableLog('rdApp.*') | |
| import yaml | |
| cache_name = datetime.now().strftime('date%d-%m_time%H-%M-%S.%f') | |
| parser = ArgumentParser() | |
| # '~/diffScreen/workdir/mdn_model_ns_48_nv_10_layer_62023-07-04_04-16-12' | |
| parser.add_argument('--config', type=FileType(mode='r'), default=None) | |
| parser.add_argument('--data_csv', type=str, default='~/DeepLearningForDock/dataset/DEKOIS2.csv', help='Path to folder with dataset for score in place') | |
| # parser.add_argument('--model_dir', type=str, default=None, help='Path to folder with trained score model and hyperparameters') | |
| parser.add_argument('--esm_embeddings_path', type=str, default="~/esm2_3billion_pdbbind_embeddings.pt", help='Path to folder with esm embeddings for screen proteins') | |
| # parser.add_argument('--ckpt', type=str, default=None, help='Checkpoint to use inside the folder') | |
| parser.add_argument('--confidence_model_dir', type=str, default=None, help='Path to folder with trained confidence model and hyperparameters') | |
| parser.add_argument('--confidence_ckpt', type=str, default=None, help='Checkpoint to use inside the folder') | |
| parser.add_argument('--model_version', type=str, default='version4', help='version of mdn model') | |
| parser.add_argument('--mdn_dist_threshold_test', type=float, default=None, help='mdn_dist_threshold_test') | |
| parser.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.') | |
| parser.add_argument('--run_name', type=str, default='test_ns_48_nv_10_layer_62023-06-25_07-54-08_model', help='') | |
| parser.add_argument('--project', type=str, default='ligbind_inf_test_mdn', help='') | |
| parser.add_argument('--out_dir', type=str, default='~/test_workdir/mdn_result_40', help='Where to save results to') | |
| parser.add_argument('--batch_size', type=int, default=40, help='Number of poses to sample in parallel') | |
| parser.add_argument('--wandb', action='store_true', default=False, help='') | |
| parser.add_argument('--wandb_dir', type=str, default='~/test_workdir', help='Folder in which to save wandb logs') | |
| parser.add_argument('--num_workers', type=int, default=1, help='Number of workers for dataset creation') | |
| args = parser.parse_args() | |
| def main_function(): | |
| """ | |
| Main function for evaluating scores in place using a confidence model. | |
| This function performs the following tasks: | |
| 1. Loads configuration settings from a YAML file if provided. | |
| 2. Sets up output directories and initializes logging. | |
| 3. Loads a pre-trained confidence model and its parameters. | |
| 4. Reads input data from a CSV file and processes it using a dataset and dataloader. | |
| 5. Evaluates the confidence model on the input data and generates predictions. | |
| 6. Saves the results to a CSV file and optionally logs them to Weights & Biases (wandb). | |
| Result CSV Columns: | |
| - `sdf_name`: The name of the SDF file associated with the molecule. | |
| - `screen_confidence(for molecule rank)`: The confidence score for ranking molecules. | |
| - `confidence_name`: The name of the confidence prediction, including detailed metadata. | |
| - `pose_prediction_confidence(for pose rank)`: The confidence score for ranking poses, extracted from `confidence_name`. | |
| - `pose_sample_idx`: The sample index of the pose, extracted from `confidence_name`. | |
| - `pose_rank`: The rank of the pose, extracted from `confidence_name`. | |
| - `molecule_name`: The name of the molecule, extracted from `confidence_name`. | |
| - `molecule_idx_in_input_file`: The index of the molecule in the input file, extracted from `confidence_name`. | |
| Notes: | |
| - The function uses the `accelerator` library for distributed processing. | |
| - The `wandb` library is used for experiment tracking if enabled. | |
| - The confidence model is loaded and evaluated in a no-gradient mode (`torch.no_grad()`). | |
| - Errors during processing of individual ligands are logged and skipped. | |
| Outputs: | |
| - A CSV file containing the results is saved in the specified output directory. | |
| - A `Readme.txt` file is generated with metadata about the run. | |
| - Optionally, results are logged to Weights & Biases (wandb). | |
| Raises: | |
| - Exceptions during ligand processing are caught and logged without halting the execution. | |
| """ | |
| if args.config: | |
| config_dict = yaml.load(args.config, Loader=yaml.FullLoader) | |
| arg_dict = args.__dict__ | |
| for key, value in config_dict.items(): | |
| if isinstance(value, list): | |
| for v in value: | |
| arg_dict[key].append(v) | |
| else: | |
| arg_dict[key] = value | |
| if args.out_dir is None: args.out_dir = f'inference_out_dir_not_specified/{args.run_name}' | |
| os.makedirs(args.out_dir, exist_ok=True) | |
| if args.confidence_model_dir is not None: | |
| with open(f'{args.confidence_model_dir}/model_parameters.yml') as f: | |
| args_dicts = yaml.full_load(f) | |
| if 'topN' not in args_dicts.keys(): | |
| args_dicts['topN'] = 1 | |
| confidence_args = Namespace(**args_dicts) | |
| # | |
| confidence_args.transfer_weights = False | |
| confidence_args.use_original_model_cache = True | |
| confidence_args.original_model_dir = None | |
| # load model param & weight bias | |
| if args.confidence_model_dir is not None: | |
| if confidence_args.transfer_weights: | |
| with open(f'{confidence_args.original_model_dir}/model_parameters.yml') as f: | |
| args_dicts = yaml.full_load(f) | |
| if 'topN' not in args_dicts.keys(): | |
| args_dicts['topN'] = 1 | |
| confidence_model_args = Namespace(**args_dicts) | |
| else: | |
| confidence_model_args = confidence_args | |
| # confidence_model_args.add_argument('--topN', type=int, default=1, help='Number of atoms to calculate confidence') | |
| confidence_model_args.mdn_dist_threshold_test = args.mdn_dist_threshold_test if args.mdn_dist_threshold_test is not None else 5.0 | |
| if not hasattr(confidence_model_args,'mdn_dist_threshold_train'): | |
| confidence_model_args.mdn_dist_threshold_train =7.0 | |
| confidence_model = get_model(confidence_model_args, device, t_to_sigma=None, no_parallel=True, | |
| model_type = 'mdn_model') | |
| state_dict = torch.load(f'{args.confidence_model_dir}/{args.confidence_ckpt}', map_location=torch.device('cpu')) | |
| confidence_model.load_state_dict(state_dict, strict=True) | |
| confidence_model = confidence_model.to(device) | |
| confidence_model.eval() | |
| confidence_model = accelerator.prepare(confidence_model) | |
| if accelerator.is_local_main_process: | |
| if args.wandb: | |
| wandb.login(key = 'yourkey') | |
| run = wandb.init( | |
| entity='SurfDock', | |
| settings=wandb.Settings(start_method="fork"), | |
| project=args.project, | |
| name=args.run_name, | |
| dir = args.wandb_dir, | |
| config=args | |
| ) | |
| df = pd.read_csv(args.data_csv) | |
| pocket_paths = df['pocket_path'].tolist() | |
| ligands_paths = df['ligand_path'].tolist() | |
| ref_ligands = df['ref_ligand'].tolist() | |
| surface_paths = df['protein_surface'].tolist() | |
| esm_embeddings_dict = torch.load(args.esm_embeddings_path) | |
| confidence = [] | |
| confidence_names = [] | |
| sdf_names = [] | |
| start_time = time.time() | |
| pbar = tqdm(zip(pocket_paths,ligands_paths,ref_ligands,surface_paths),total=len(pocket_paths)) | |
| for pocket_path,ligands_path,ref_ligand,surface_path in pbar: | |
| try: | |
| esm_embeddings = esm_embeddings_dict[os.path.splitext(os.path.basename(pocket_path))[0]] | |
| # assert False, f'esmembedding shape : {esm_embeddings.shape}' | |
| test_dataset = ScreenDataset(pocket_path,ligands_path,ref_ligand,surface_path,transform=None, | |
| receptor_radius=confidence_args.receptor_radius, | |
| cache_path=None, split_path=None, | |
| remove_hs=confidence_args.remove_hs, max_lig_size=None, | |
| c_alpha_max_neighbors=confidence_args.c_alpha_max_neighbors, | |
| matching= False, keep_original=True, | |
| popsize=confidence_args.matching_popsize, | |
| maxiter=confidence_args.matching_maxiter, | |
| all_atoms=confidence_args.all_atoms, | |
| atom_radius=confidence_args.atom_radius, | |
| atom_max_neighbors=confidence_args.atom_max_neighbors, | |
| esm_embeddings=esm_embeddings, | |
| require_ligand=False, | |
| num_workers=args.num_workers) | |
| test_loader = DataLoader(dataset=test_dataset, batch_size=args.batch_size, shuffle=False,num_workers=args.num_workers) | |
| if len(test_dataset) == 0: | |
| continue | |
| test_loader= accelerator.prepare(test_loader) | |
| logger.info('Size of test dataset: ', len(test_dataset)) | |
| with torch.no_grad(): | |
| confidence_model.eval() | |
| for confidence_complex_graph_batch in tqdm(test_loader,total = len(test_loader)): | |
| confidence += confidence_model(confidence_complex_graph_batch)[-1].cpu().detach().numpy().tolist() | |
| confidence_names += confidence_complex_graph_batch['name'] | |
| sdf_names += [os.path.basename(ligands_path)]*len(confidence_complex_graph_batch['name']) | |
| assert len(confidence)==len(confidence_names)==len(sdf_names) | |
| # logger.info(len(confidence_complex_graph_batch['name'][0]),len(confidence_complex_graph_batch['name']),confidence_complex_graph_batch['name'][0]) | |
| except Exception as e: | |
| logger.info(e,'some error failed for : ',ligands_path) | |
| continue | |
| # if accelerator.is_local_main_process: | |
| pbar.set_description('screen time used: {:.2f} '.format(time.time()-start_time)) | |
| logger.info('screen time used: ',time.time()-start_time) | |
| if accelerator.is_local_main_process: | |
| result = pd.DataFrame({'sdf_name':sdf_names,'screen_confidence(for molecule rank)':confidence,'pose_file_path':confidence_names}) | |
| csv_flag = os.path.basename(args.data_csv).split('.')[0] | |
| result['pose_prediction_confidence(for pose rank)'] = result['pose_file_path'].apply(lambda x: float(x.split('_')[-1].split('.sdf')[0])) | |
| result['pose_sample_idx'] = result['pose_file_path'].apply(lambda x: float(x.split('sample_idx_')[-1].split('_rank')[0])) | |
| result['pose_rank'] = result['pose_file_path'].apply(lambda x: float(x.split('_rank_')[-1].split('_confidence_')[0])) | |
| result['molecule_name'] = result['pose_file_path'].apply(lambda x: x.split('.sdf_file_inner_idx_')[0].split('/')[-1]) | |
| result['molecule_idx_in_input_file'] = result['pose_file_path'].apply(lambda x: float(x.split('sdf_file_inner_idx_')[-1].split('_sample_idx_')[0])) | |
| result.to_csv(f'{args.out_dir}/{csv_flag}_confidence.csv',index=False) | |
| with open(f"{args.out_dir}/Readme.txt", 'w') as f: | |
| f.write(f"""Result CSV Columns: | |
| - `sdf_name`: The name of the SDF file associated with the molecule. | |
| - `screen_confidence(for molecule rank)`: The confidence score for ranking molecules to screen a library. | |
| - `pose_file_path`: The path of the SDF file associated with the molecule, including detailed metadata. | |
| - `pose_prediction_confidence(for pose rank)`: The confidence score for ranking poses, extracted from `pose_file_path`. | |
| - `pose_sample_idx`: The sample index of the pose, extracted from `pose_file_path`. | |
| - `pose_rank`: The rank of the pose, extracted from `pose_file_path`. | |
| - `molecule_name`: The name of the molecule, extracted from `pose_file_path`. | |
| - `molecule_idx_in_input_file`: The index of the molecule in the input file, extracted from `pose_file_path`. | |
| """) | |
| # np.save(f'{args.out_dir}/confidence.npy', confidence) | |
| if args.wandb: | |
| wandb.finish() | |
| if __name__ == '__main__': | |
| from accelerate import Accelerator | |
| from accelerate.utils import DistributedDataParallelKwargs | |
| kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) | |
| accelerator = Accelerator(kwargs_handlers=[kwargs]) | |
| from accelerate.utils import set_seed | |
| import sys | |
| device = accelerator.device | |
| set_seed(42) | |
| accelerator.print(f'device {str(accelerator.device)} is used!') | |
| main_function() | |
| sys.exit() | |