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()