File size: 13,563 Bytes
96f168d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 | """
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()
|