File size: 4,925 Bytes
3ac1d94
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import argparse
import multiprocessing as mp
import os
import pickle
from functools import partial

from rdkit import RDLogger
from tqdm.auto import tqdm

from onescience.datapipes.targetdiff.protein_ligand import parse_sdf_file_mol, read_mol, KMAP
from onescience.utils.targetdiff.data import PDBProtein

RDLogger.DisableLog('rdApp.*')


def parse_pdbbind_index_file(raw_path, subset='refined'):
    all_index = []
    version = int(raw_path.rstrip('/')[-4:])
    assert version >= 2016
    if subset == 'refined':
        data_path = os.path.join(raw_path, f'refined-set')
        index_path = os.path.join(data_path, 'index', f'INDEX_refined_data.{version}')
    elif subset == 'general':
        data_path = os.path.join(raw_path, f'general-set-except-refined')
        index_path = os.path.join(data_path, 'index', f'INDEX_general_PL_data.{version}')
    else:
        raise ValueError(subset)

    all_files = os.listdir(data_path)
    with open(index_path, 'r') as f:
        lines = f.readlines()
    for line in lines:
        if line.startswith('#'): continue
        index, res, year, pka, kv = line.split('//')[0].strip().split()
        kind = [v for k, v in KMAP.items() if k in kv]
        assert len(kind) == 1
        if index in all_files:
            all_index.append([index, res, year, pka, kind[0]])
    return all_index


def process_item(item, args):
    pdb_idx, res, year, pka, kind = item
    ligand_path = None
    try:
        if args.subset == 'refined':
            pdb_path = os.path.join(args.source, 'refined-set', pdb_idx)
        elif args.subset == 'general':
            pdb_path = os.path.join(args.source, 'general-set-except-refined', pdb_idx)
        else:
            raise ValueError(args.subset)

        protein_path = os.path.join(pdb_path, f'{pdb_idx}_protein.pdb')
        ligand_sdf_path = os.path.join(pdb_path, f'{pdb_idx}_ligand.sdf')
        ligand_mol2_path = os.path.join(pdb_path, f'{pdb_idx}_ligand.mol2')
        mol, problem, ligand_path = read_mol(ligand_sdf_path, ligand_mol2_path)
        if problem:
            print('Read mol error.', item)
            return None, ligand_path, res, pka, kind

        protein = PDBProtein(protein_path)
        # ligand = parse_sdf_file_mol(ligand_path, heavy_only=True)
        ligand = parse_sdf_file_mol(ligand_path, heavy_only=False)
        if args.dest is None:
            pocket_path = os.path.join(pdb_path, f'{pdb_idx}_pocket{args.radius}.pdb')
        else:
            pocket_dir = os.path.join(args.dest, f'pocket_{args.radius}_{args.subset}', 'pockets')
            os.makedirs(pocket_dir, exist_ok=True)
            pocket_path = os.path.join(pocket_dir, f'{pdb_idx}_pocket{args.radius}.pdb')
        if not os.path.exists(pocket_path):
            pdb_block_pocket = protein.residues_to_pdb_block(
                protein.query_residues_ligand(ligand, args.radius)
            )
            with open(pocket_path, 'w') as f:
                f.write(pdb_block_pocket)
        return pocket_path, ligand_path, res, pka, kind

    except Exception as e:
        print('Exception occured.', item, repr(e))
        return None, ligand_path, res, pka, kind


if __name__ == '__main__':
    parser = argparse.ArgumentParser()
    parser.add_argument('--source', type=str, default='./data/pdbbind_v2016')
    parser.add_argument('--dest', type=str, default=None)
    parser.add_argument('--fixed_sdf_dir', type=str, default='./data/pdbbind_v2016/fixed_sdf_files')
    parser.add_argument('--subset', type=str, default='refined')
    parser.add_argument('--refined_index_pkl', type=str, default=None)
    parser.add_argument('--radius', type=int, default=10)
    parser.add_argument('--num_workers', type=int, default=16)
    args = parser.parse_args()

    index = parse_pdbbind_index_file(args.source, args.subset)
    # if not os.path.exists(args.fixed_sdf_dir):
    #     os.makedirs(args.fixed_sdf_dir)

    pool = mp.Pool(args.num_workers)
    index_pocket = []
    for item_pocket in tqdm(pool.imap_unordered(partial(process_item, args=args), index), total=len(index)):
        index_pocket.append(item_pocket)
    pool.close()

    valid_index_pocket = []
    for index in index_pocket:
        if index[0] is not None:
            valid_index_pocket.append(index)

    output_root = args.source if args.dest is None else args.dest
    save_path = os.path.join(output_root, f'pocket_{args.radius}_{args.subset}')
    if not os.path.exists(save_path):
        os.makedirs(save_path)
    index_path = os.path.join(save_path, 'index.pkl')

    if args.subset == 'general' and args.refined_index_pkl is not None:
        with open(args.refined_index_pkl, 'rb') as f:
            refined_index = pickle.load(f)
        valid_index_pocket += refined_index
    with open(index_path, 'wb') as f:
        pickle.dump(valid_index_pocket, f)
    print('Done. %d protein-ligand pairs in total.' % len(valid_index_pocket))