| |
| |
| |
| import json |
|
|
| |
| import logging |
| import os |
| |
| |
| from xmlrpc.client import Boolean |
| import numpy as np |
| import torch |
| import pickle |
| from tqdm import tqdm |
| from unicore import checkpoint_utils |
| import unicore |
| from unicore.data import (AppendTokenDataset, Dictionary, EpochShuffleDataset, |
| FromNumpyDataset, NestedDictionaryDataset, |
| PrependTokenDataset, RawArrayDataset, LMDBDataset, RawLabelDataset, |
| RightPadDataset, RightPadDataset2D, TokenizeDataset, SortDataset, data_utils) |
| from unicore.tasks import UnicoreTask, register_task |
| from unimol.data import (AffinityDataset, CroppingPocketDataset, |
| CrossDistanceDataset, DistanceDataset, |
| EdgeTypeDataset, KeyDataset, LengthDataset, |
| NormalizeDataset, NormalizeDockingPoseDataset, |
| PrependAndAppend2DDataset, RemoveHydrogenDataset, |
| RemoveHydrogenPocketDataset, RightPadDatasetCoord, |
| RightPadDatasetCross2D, TTADockingPoseDataset, AffinityTestDataset, AffinityValidDataset, |
| AffinityMolDataset, AffinityPocketDataset, ResamplingDataset) |
| |
| from rdkit.ML.Scoring.Scoring import CalcBEDROC, CalcAUC, CalcEnrichment |
| from sklearn.metrics import roc_curve |
|
|
| logger = logging.getLogger(__name__) |
| import os |
|
|
| PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) |
|
|
|
|
| def re_new(y_true, y_score, ratio): |
| fp = 0 |
| tp = 0 |
| p = sum(y_true) |
| n = len(y_true) - p |
| num = ratio * n |
| sort_index = np.argsort(y_score)[::-1] |
| for i in range(len(sort_index)): |
| index = sort_index[i] |
| if y_true[index] == 1: |
| tp += 1 |
| else: |
| fp += 1 |
| if fp >= num: |
| break |
| return (tp * n) / (p * fp) |
|
|
|
|
| def calc_re(y_true, y_score, ratio_list): |
| res2 = {} |
|
|
| for ratio in ratio_list: |
| res2[str(ratio)] = re_new(y_true, y_score, ratio) |
|
|
| return res2 |
|
|
|
|
| def cal_metrics(y_true, y_score, alpha): |
| """ |
| Calculate BEDROC score. |
| |
| Parameters: |
| - y_true: true binary labels (0 or 1) |
| - y_score: predicted scores or probabilities |
| - alpha: parameter controlling the degree of early retrieval emphasis |
| |
| Returns: |
| - BEDROC score |
| """ |
|
|
| |
| scores = np.expand_dims(y_score, axis=1) |
| y_true = np.expand_dims(y_true, axis=1) |
| scores = np.concatenate((scores, y_true), axis=1) |
| |
| scores = scores[scores[:, 0].argsort()[::-1]] |
| bedroc = CalcBEDROC(scores, 1, 80.5) |
| count = 0 |
| |
| index = np.argsort(y_score)[::-1] |
| for i in range(int(len(index) * 0.005)): |
| if y_true[index[i]] == 1: |
| count += 1 |
| auc = CalcAUC(scores, 1) |
| ef_list = CalcEnrichment(scores, 1, [0.005, 0.01, 0.02, 0.05]) |
| ef = { |
| "0.005": ef_list[0], |
| "0.01": ef_list[1], |
| "0.02": ef_list[2], |
| "0.05": ef_list[3] |
| } |
| re_list = calc_re(y_true, y_score, [0.005, 0.01, 0.02, 0.05]) |
| return auc, bedroc, ef, re_list |
|
|
|
|
| def get_uniprot_seq(uniprot): |
| import urllib |
| if not os.path.exists(f"./uniprot_fasta/{uniprot}.fasta"): |
| os.system(f"mkdir -p ./uniprot_fasta") |
| urllib.request.urlretrieve(f"https://rest.uniprot.org/uniprotkb/{uniprot}.fasta", |
| f"./uniprot_fasta/{uniprot}.fasta") |
|
|
| with open(f"./uniprot_fasta/{uniprot}.fasta", "r") as f: |
| lines = [] |
| for line in f.readlines(): |
| if line.startswith(">"): |
| continue |
| else: |
| lines.append(line.strip()) |
| return "".join(lines) |
|
|
|
|
| @register_task("test_task") |
| class ContrasRankTest(UnicoreTask): |
| """Task for training transformer auto-encoder models.""" |
|
|
| @staticmethod |
| def add_args(parser): |
| """Add task-specific arguments to the parser.""" |
| parser.add_argument( |
| "data", |
| help="downstream data path", |
| ) |
| parser.add_argument( |
| "--finetune-mol-model", |
| default=None, |
| type=str, |
| help="pretrained molecular model path", |
| ) |
| parser.add_argument( |
| "--finetune-pocket-model", |
| default=None, |
| type=str, |
| help="pretrained pocket model path", |
| ) |
| parser.add_argument( |
| "--dist-threshold", |
| type=float, |
| default=6.0, |
| help="threshold for the distance between the molecule and the pocket", |
| ) |
| parser.add_argument( |
| "--max-pocket-atoms", |
| type=int, |
| default=256, |
| help="selected maximum number of atoms in a pocket", |
| ) |
| parser.add_argument( |
| "--test-model", |
| default=False, |
| type=Boolean, |
| help="whether test model", |
| ) |
| parser.add_argument( |
| "--demo-lig-file", |
| type=str, |
| default="" |
| ) |
| parser.add_argument( |
| "--demo-prot-file", |
| type=str, |
| default="" |
| ) |
| parser.add_argument( |
| "--demo-uniprot", |
| type=str, |
| default="" |
| ) |
| parser.add_argument("--reg", action="store_true", help="regression task") |
|
|
| def __init__(self, args, dictionary, pocket_dictionary): |
| super().__init__(args) |
| self.dictionary = dictionary |
| self.pocket_dictionary = pocket_dictionary |
| self.seed = args.seed |
| |
| self.mask_idx = dictionary.add_symbol("[MASK]", is_special=True) |
| self.pocket_mask_idx = pocket_dictionary.add_symbol("[MASK]", is_special=True) |
| self.mol_reps = None |
| self.keys = None |
|
|
| @classmethod |
| def setup_task(cls, args, **kwargs): |
| mol_dictionary = Dictionary.load(os.path.join(PROJECT_ROOT, "vocab", "dict_mol.txt")) |
| pocket_dictionary = Dictionary.load(os.path.join(PROJECT_ROOT, "vocab", "dict_pkt.txt")) |
| logger.info("ligand dictionary: {} types".format(len(mol_dictionary))) |
| logger.info("pocket dictionary: {} types".format(len(pocket_dictionary))) |
| return cls(args, mol_dictionary, pocket_dictionary) |
|
|
| def load_dataset(self, split, **kwargs): |
| """Load a given dataset split. |
| 'smi','pocket','atoms','coordinates','pocket_atoms','pocket_coordinates' |
| Args: |
| split (str): name of the data scoure (e.g., bppp) |
| """ |
| if split == "test": |
| data_path = f"{self.args.data}/casf.lmdb" |
| else: |
| data_path = os.path.join(self.args.data, split + ".lmdb") |
| dataset = LMDBDataset(data_path) |
| if split.startswith("train"): |
| smi_dataset = KeyDataset(dataset, "smi") |
| poc_dataset = KeyDataset(dataset, "pocket") |
|
|
| dataset = AffinityDataset( |
| dataset, |
| self.args.seed, |
| "atoms", |
| "coordinates", |
| "pocket_atoms", |
| "pocket_coordinates", |
| "label", |
| True, |
| ) |
| tgt_dataset = KeyDataset(dataset, "affinity") |
|
|
| else: |
|
|
| dataset = AffinityDataset( |
| dataset, |
| self.args.seed, |
| "atoms", |
| "coordinates", |
| "pocket_atoms", |
| "pocket_coordinates", |
| "label", |
| ) |
| tgt_dataset = KeyDataset(dataset, "affinity") |
| smi_dataset = KeyDataset(dataset, "smi") |
| poc_dataset = KeyDataset(dataset, "pocket") |
|
|
| def PrependAndAppend(dataset, pre_token, app_token): |
| dataset = PrependTokenDataset(dataset, pre_token) |
| return AppendTokenDataset(dataset, app_token) |
|
|
| dataset = RemoveHydrogenPocketDataset( |
| dataset, |
| "pocket_atoms", |
| "pocket_coordinates", |
| True, |
| True, |
| ) |
| dataset = CroppingPocketDataset( |
| dataset, |
| self.seed, |
| "pocket_atoms", |
| "pocket_coordinates", |
| self.args.max_pocket_atoms, |
| ) |
|
|
| dataset = RemoveHydrogenDataset(dataset, "atoms", "coordinates", True, True) |
|
|
| apo_dataset = NormalizeDataset(dataset, "coordinates") |
| apo_dataset = NormalizeDataset(apo_dataset, "pocket_coordinates") |
|
|
| src_dataset = KeyDataset(apo_dataset, "atoms") |
| mol_len_dataset = LengthDataset(src_dataset) |
| src_dataset = TokenizeDataset( |
| src_dataset, self.dictionary, max_seq_len=self.args.max_seq_len |
| ) |
| coord_dataset = KeyDataset(apo_dataset, "coordinates") |
| src_dataset = PrependAndAppend( |
| src_dataset, self.dictionary.bos(), self.dictionary.eos() |
| ) |
| edge_type = EdgeTypeDataset(src_dataset, len(self.dictionary)) |
| coord_dataset = FromNumpyDataset(coord_dataset) |
| distance_dataset = DistanceDataset(coord_dataset) |
| coord_dataset = PrependAndAppend(coord_dataset, 0.0, 0.0) |
| distance_dataset = PrependAndAppend2DDataset(distance_dataset, 0.0) |
|
|
| src_pocket_dataset = KeyDataset(apo_dataset, "pocket_atoms") |
| pocket_len_dataset = LengthDataset(src_pocket_dataset) |
| src_pocket_dataset = TokenizeDataset( |
| src_pocket_dataset, |
| self.pocket_dictionary, |
| max_seq_len=self.args.max_seq_len, |
| ) |
| coord_pocket_dataset = KeyDataset(apo_dataset, "pocket_coordinates") |
| src_pocket_dataset = PrependAndAppend( |
| src_pocket_dataset, |
| self.pocket_dictionary.bos(), |
| self.pocket_dictionary.eos(), |
| ) |
| pocket_edge_type = EdgeTypeDataset( |
| src_pocket_dataset, len(self.pocket_dictionary) |
| ) |
| coord_pocket_dataset = FromNumpyDataset(coord_pocket_dataset) |
| distance_pocket_dataset = DistanceDataset(coord_pocket_dataset) |
| coord_pocket_dataset = PrependAndAppend(coord_pocket_dataset, 0.0, 0.0) |
| distance_pocket_dataset = PrependAndAppend2DDataset( |
| distance_pocket_dataset, 0.0 |
| ) |
|
|
| nest_dataset = NestedDictionaryDataset( |
| { |
| "net_input": { |
| "mol_src_tokens": RightPadDataset( |
| src_dataset, |
| pad_idx=self.dictionary.pad(), |
| ), |
| "mol_src_distance": RightPadDataset2D( |
| distance_dataset, |
| pad_idx=0, |
| ), |
| "mol_src_edge_type": RightPadDataset2D( |
| edge_type, |
| pad_idx=0, |
| ), |
| "pocket_src_tokens": RightPadDataset( |
| src_pocket_dataset, |
| pad_idx=self.pocket_dictionary.pad(), |
| ), |
| "pocket_src_distance": RightPadDataset2D( |
| distance_pocket_dataset, |
| pad_idx=0, |
| ), |
| "pocket_src_edge_type": RightPadDataset2D( |
| pocket_edge_type, |
| pad_idx=0, |
| ), |
| "pocket_src_coord": RightPadDatasetCoord( |
| coord_pocket_dataset, |
| pad_idx=0, |
| ), |
| "mol_len": RawArrayDataset(mol_len_dataset), |
| "pocket_len": RawArrayDataset(pocket_len_dataset) |
| }, |
| "target": { |
| "finetune_target": RawLabelDataset(tgt_dataset), |
| }, |
| "smi_name": RawArrayDataset(smi_dataset), |
| "pocket_name": RawArrayDataset(poc_dataset), |
| }, |
| ) |
| if split == "train" and kwargs.get("shuffle", True): |
| with data_utils.numpy_seed(self.args.seed): |
| shuffle = np.random.permutation(len(src_dataset)) |
|
|
| self.datasets[split] = SortDataset( |
| nest_dataset, |
| sort_order=[shuffle], |
| ) |
| self.datasets[split] = ResamplingDataset( |
| self.datasets[split] |
| ) |
| else: |
| self.datasets[split] = nest_dataset |
|
|
| return self.datasets[split] |
|
|
| def load_mols_dataset(self, data_path, atoms, coords, **kwargs): |
|
|
| dataset = LMDBDataset(data_path) |
| |
| try: |
| label_dataset = KeyDataset(dataset, "label") |
| x = label_dataset[0] |
| except: |
| label_dataset = None |
| dataset = AffinityMolDataset( |
| dataset, |
| self.args.seed, |
| atoms, |
| coords, |
| False, |
| ) |
|
|
| smi_dataset = KeyDataset(dataset, "smi") |
| mol_dataset = KeyDataset(dataset, "mol") |
| if kwargs.get("load_name", False): |
| name_dataset = KeyDataset(dataset, "name") |
|
|
| def PrependAndAppend(dataset, pre_token, app_token): |
| dataset = PrependTokenDataset(dataset, pre_token) |
| return AppendTokenDataset(dataset, app_token) |
|
|
| dataset = RemoveHydrogenDataset(dataset, "atoms", "coordinates", True, True) |
| apo_dataset = NormalizeDataset(dataset, "coordinates") |
|
|
| src_dataset = KeyDataset(apo_dataset, "atoms") |
| len_dataset = LengthDataset(src_dataset) |
| src_dataset = TokenizeDataset( |
| src_dataset, self.dictionary, max_seq_len=self.args.max_seq_len |
| ) |
| coord_dataset = KeyDataset(apo_dataset, "coordinates") |
| src_dataset = PrependAndAppend( |
| src_dataset, self.dictionary.bos(), self.dictionary.eos() |
| ) |
| edge_type = EdgeTypeDataset(src_dataset, len(self.dictionary)) |
| coord_dataset = FromNumpyDataset(coord_dataset) |
| distance_dataset = DistanceDataset(coord_dataset) |
| coord_dataset = PrependAndAppend(coord_dataset, 0.0, 0.0) |
| distance_dataset = PrependAndAppend2DDataset(distance_dataset, 0.0) |
|
|
| if label_dataset is not None: |
| in_datasets = { |
| "net_input": { |
| "mol_src_tokens": RightPadDataset( |
| src_dataset, |
| pad_idx=self.dictionary.pad(), |
| ), |
| "mol_src_distance": RightPadDataset2D( |
| distance_dataset, |
| pad_idx=0, |
| ), |
| "mol_src_edge_type": RightPadDataset2D( |
| edge_type, |
| pad_idx=0, |
| ), |
| }, |
| "smi_name": RawArrayDataset(smi_dataset), |
| "target": RawArrayDataset(label_dataset), |
| "mol_len": RawArrayDataset(len_dataset), |
| "mol": RawArrayDataset(mol_dataset) |
| } |
| else: |
| in_datasets = { |
| "net_input": { |
| "mol_src_tokens": RightPadDataset( |
| src_dataset, |
| pad_idx=self.dictionary.pad(), |
| ), |
| "mol_src_distance": RightPadDataset2D( |
| distance_dataset, |
| pad_idx=0, |
| ), |
| "mol_src_edge_type": RightPadDataset2D( |
| edge_type, |
| pad_idx=0, |
| ), |
| }, |
| "smi_name": RawArrayDataset(smi_dataset), |
| |
| "mol_len": RawArrayDataset(len_dataset), |
| "mol": RawArrayDataset(mol_dataset) |
| } |
| if kwargs.get("load_name", False): |
| in_datasets["name"] = name_dataset |
|
|
| nest_dataset = NestedDictionaryDataset(in_datasets) |
| return nest_dataset |
|
|
| def load_pockets_dataset(self, data_path, **kwargs): |
|
|
| dataset = LMDBDataset(data_path) |
|
|
| dataset = AffinityPocketDataset( |
| dataset, |
| self.args.seed, |
| "pocket_atoms", |
| "pocket_coordinates", |
| False, |
| "pocket" |
| ) |
| poc_dataset = KeyDataset(dataset, "pocket") |
| resname_dataset = KeyDataset(dataset, "pocket_residue_name") |
|
|
| def PrependAndAppend(dataset, pre_token, app_token): |
| dataset = PrependTokenDataset(dataset, pre_token) |
| return AppendTokenDataset(dataset, app_token) |
|
|
| dataset = RemoveHydrogenPocketDataset( |
| dataset, |
| "pocket_atoms", |
| "pocket_coordinates", |
| True, |
| True, |
| ) |
| dataset = CroppingPocketDataset( |
| dataset, |
| self.seed, |
| "pocket_atoms", |
| "pocket_coordinates", |
| self.args.max_pocket_atoms, |
| ) |
|
|
| apo_dataset = NormalizeDataset(dataset, "pocket_coordinates") |
|
|
| src_pocket_dataset = KeyDataset(apo_dataset, "pocket_atoms") |
| len_dataset = LengthDataset(src_pocket_dataset) |
| src_pocket_dataset = TokenizeDataset( |
| src_pocket_dataset, |
| self.pocket_dictionary, |
| max_seq_len=self.args.max_seq_len, |
| ) |
| coord_pocket_dataset = KeyDataset(apo_dataset, "pocket_coordinates") |
| src_pocket_dataset = PrependAndAppend( |
| src_pocket_dataset, |
| self.pocket_dictionary.bos(), |
| self.pocket_dictionary.eos(), |
| ) |
| pocket_edge_type = EdgeTypeDataset( |
| src_pocket_dataset, len(self.pocket_dictionary) |
| ) |
| coord_pocket_dataset = FromNumpyDataset(coord_pocket_dataset) |
| distance_pocket_dataset = DistanceDataset(coord_pocket_dataset) |
| coord_pocket_dataset = PrependAndAppend(coord_pocket_dataset, 0.0, 0.0) |
| distance_pocket_dataset = PrependAndAppend2DDataset( |
| distance_pocket_dataset, 0.0 |
| ) |
|
|
| nest_dataset = NestedDictionaryDataset( |
| { |
| "net_input": { |
| "pocket_src_tokens": RightPadDataset( |
| src_pocket_dataset, |
| pad_idx=self.pocket_dictionary.pad(), |
| ), |
| "pocket_src_distance": RightPadDataset2D( |
| distance_pocket_dataset, |
| pad_idx=0, |
| ), |
| "pocket_src_edge_type": RightPadDataset2D( |
| pocket_edge_type, |
| pad_idx=0, |
| ), |
| "pocket_src_coord": RightPadDatasetCoord( |
| coord_pocket_dataset, |
| pad_idx=0, |
| ), |
| }, |
| "pocket_residue_names": RawArrayDataset(resname_dataset), |
| "pocket_name": RawArrayDataset(poc_dataset), |
| "pocket_len": RawArrayDataset(len_dataset), |
| }, |
| ) |
| return nest_dataset |
|
|
| def build_model(self, args): |
| from unicore import models |
|
|
| model = models.build_model(args, self) |
|
|
| if args.finetune_mol_model is not None: |
| print("load pretrain model weight from...", args.finetune_mol_model) |
| state = checkpoint_utils.load_checkpoint_to_cpu( |
| args.finetune_mol_model, |
| ) |
| model.mol_model.load_state_dict(state["model"], strict=False) |
|
|
| if args.finetune_pocket_model is not None: |
| print("load pretrain model weight from...", args.finetune_pocket_model) |
| state = checkpoint_utils.load_checkpoint_to_cpu( |
| args.finetune_pocket_model, |
| ) |
| model.pocket_model.load_state_dict(state["model"], strict=False) |
|
|
| return model |
|
|
| def test_pcba_target(self, name, model, seq, **kwargs): |
| """Encode a dataset with the molecule encoder.""" |
|
|
| |
| data_path = f"{self.args.data}/lit_pcba/" + name + "/mols.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = self.args.batch_size |
| |
| mol_reps = [] |
| mol_names = [] |
| labels = [] |
|
|
| |
|
|
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, num_workers=8, |
| collate_fn=mol_dataset.collater) |
| for _, sample in enumerate(tqdm(mol_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| mol_emb = model.mol_forward(**sample["net_input"]) |
| mol_emb = mol_emb.detach().cpu().numpy() |
| mol_reps.append(mol_emb) |
| mol_names.extend(sample["smi_name"]) |
| labels.extend(sample["target"].detach().cpu().numpy()) |
| mol_reps = np.concatenate(mol_reps, axis=0) |
| labels = np.array(labels, dtype=np.int32) |
| |
| data_path = f"{self.args.data}/lit_pcba/" + name + "/pockets.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, collate_fn=pocket_dataset.collater) |
| pocket_reps = [] |
| pocket_names = [] |
|
|
| for _, sample in enumerate(tqdm(pocket_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| pocket_emb = model.pocket_forward(protein_sequences=seq, **sample["net_input"]) |
| pocket_emb = pocket_emb.detach().cpu().numpy() |
| pocket_name = sample["pocket_name"] |
| pocket_names.append(pocket_name) |
| pocket_reps.append(pocket_emb) |
| pocket_reps = np.concatenate(pocket_reps, axis=0) |
|
|
| os.system(f"mkdir -p {self.args.results_path}/PCBA/{name}") |
| np.save(f"{self.args.results_path}/PCBA/{name}/saved_mols_embed.npy", mol_reps) |
| np.save(f"{self.args.results_path}/PCBA/{name}/saved_target_embed.npy", pocket_reps) |
| np.save(f"{self.args.results_path}/PCBA/{name}/saved_labels.npy", labels) |
| json.dump(pocket_names, open(f"{self.args.results_path}/PCBA/{name}/saved_pocket_names.json", "w")) |
|
|
| res = pocket_reps @ mol_reps.T |
| res_single = res.max(axis=0) |
| auc, bedroc, ef_list, re_list = cal_metrics(labels, res_single, 80.5) |
|
|
| return auc, bedroc, ef_list, re_list |
|
|
| def test_pcba_target_regression(self, name, model, seq, **kwargs): |
| """Encode a dataset with the molecule encoder.""" |
|
|
| |
| data_path = f"{self.args.data}/lit_pcba/" + name + "/mols.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = self.args.batch_size |
| |
| mol_names = [] |
| labels = [] |
| act_preds_all = [] |
|
|
| |
|
|
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, collate_fn=mol_dataset.collater, |
| num_workers=8) |
| for _, mol_sample in enumerate(tqdm(mol_data)): |
| mol_sample = unicore.utils.move_to_cuda(mol_sample) |
| mol_names.extend(mol_sample["smi_name"]) |
| labels.extend(mol_sample["target"].detach().cpu().numpy()) |
|
|
| |
| data_path = f"{self.args.data}/lit_pcba/" + name + "/pockets.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, |
| collate_fn=pocket_dataset.collater) |
| act_preds = [] |
| pocket_names = [] |
|
|
| for _, pocket_sample in enumerate(pocket_data): |
| pocket_sample = unicore.utils.move_to_cuda(pocket_sample) |
| pred = model.forward(protein_sequences=seq, **pocket_sample["net_input"], **mol_sample["net_input"]) |
| pocket_name = pocket_sample["pocket_name"] |
| act_preds.append(pred.detach().cpu().numpy()) |
| pocket_names.append(pocket_name) |
|
|
| act_preds = np.concatenate(act_preds, axis=0) |
| act_preds_all.append(act_preds) |
|
|
| labels = np.array(labels, dtype=np.int32) |
| res = np.concatenate(act_preds_all, axis=1) |
| res_single = res.max(axis=0) |
| os.system(f"mkdir -p {self.args.results_path}/PCBA/{name}") |
| np.save(f"{self.args.results_path}/PCBA/{name}/saved_labels.npy", labels) |
| np.save(f"{self.args.results_path}/PCBA/{name}/saved_preds.npy", res_single) |
| json.dump(pocket_names, open(f"{self.args.results_path}/PCBA/{name}/saved_pocket_names.json", "w")) |
|
|
| auc, bedroc, ef_list, re_list = cal_metrics(labels, res_single, 80.5) |
|
|
| return auc, bedroc, ef_list, re_list |
|
|
| def test_pcba(self, model, **kwargs): |
| targets = os.listdir(f"{self.args.data}/lit_pcba/") |
|
|
| |
| auc_list = [] |
| ef_list = [] |
| bedroc_list = [] |
|
|
| re_list = { |
| "0.005": [], |
| "0.01": [], |
| "0.02": [], |
| "0.05": [] |
| } |
| ef_list = { |
| "0.005": [], |
| "0.01": [], |
| "0.02": [], |
| "0.05": [] |
| } |
| uniprot_list = json.load(open(f"{self.args.data}/PCBA.json")) |
| target2uniport = {x[2]: x[0] for x in uniprot_list} |
|
|
| for target in targets: |
| print(target) |
| |
| |
| seq = get_uniprot_seq(target2uniport[target]) |
| if self.args.arch in ["DTA", "pocketregression"]: |
| auc, bedroc, ef, re = self.test_pcba_target_regression(target, model, seq) |
| else: |
| auc, bedroc, ef, re = self.test_pcba_target(target, model, seq) |
| auc_list.append(auc) |
| bedroc_list.append(bedroc) |
| for key in ef: |
| ef_list[key].append(ef[key]) |
| print("re", re) |
| print("ef", ef) |
| for key in re: |
| re_list[key].append(re[key]) |
| print(auc_list) |
| print(ef_list) |
| print("auc 25%", np.percentile(auc_list, 25)) |
| print("auc 50%", np.percentile(auc_list, 50)) |
| print("auc 75%", np.percentile(auc_list, 75)) |
| print("auc mean", np.mean(auc_list)) |
| print("bedroc 25%", np.percentile(bedroc_list, 25)) |
| print("bedroc 50%", np.percentile(bedroc_list, 50)) |
| print("bedroc 75%", np.percentile(bedroc_list, 75)) |
| print("bedroc mean", np.mean(bedroc_list)) |
| |
| |
| for key in ef_list: |
| print("ef", key, "25%", np.percentile(ef_list[key], 25)) |
| print("ef", key, "50%", np.percentile(ef_list[key], 50)) |
| print("ef", key, "75%", np.percentile(ef_list[key], 75)) |
| print("ef", key, "mean", np.mean(ef_list[key])) |
| for key in re_list: |
| print("re", key, "25%", np.percentile(re_list[key], 25)) |
| print("re", key, "50%", np.percentile(re_list[key], 50)) |
| print("re", key, "75%", np.percentile(re_list[key], 75)) |
| print("re", key, "mean", np.mean(re_list[key])) |
|
|
| return |
|
|
| def test_dude_target(self, target, model, seq, **kwargs): |
|
|
| data_path = f"{self.args.data}/DUD-E/" + target + "/mols_real.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = 64 |
| print(num_data // bsz) |
| mol_reps = [] |
| mol_names = [] |
| labels = [] |
|
|
| |
| print("begin with target:", target) |
| print("number of mol:", len(mol_dataset)) |
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, num_workers=8, |
| collate_fn=mol_dataset.collater) |
|
|
| for _, sample in enumerate(tqdm(mol_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| mol_emb = model.mol_forward(**sample["net_input"]) |
| mol_emb = mol_emb.detach().cpu().numpy() |
| |
| mol_reps.append(mol_emb) |
| mol_names.extend(sample["smi_name"]) |
| labels.extend(sample["target"].detach().cpu().numpy()) |
| mol_reps = np.concatenate(mol_reps, axis=0) |
| labels = np.array(labels, dtype=np.int32) |
| |
| data_path = f"{self.args.data}/DUD-E/" + target + "/pocket.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, collate_fn=pocket_dataset.collater) |
| pocket_reps = [] |
|
|
| for _, sample in enumerate(tqdm(pocket_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| pocket_emb = model.pocket_forward(protein_sequences=seq, **sample["net_input"]) |
| pocket_emb = pocket_emb.detach().cpu().numpy() |
| pocket_reps.append(pocket_emb) |
| pocket_reps = np.concatenate(pocket_reps, axis=0) |
| print(pocket_reps.shape) |
| res = pocket_reps @ mol_reps.T |
|
|
| res_single = res.max(axis=0) |
| os.system(f"mkdir -p {self.args.results_path}/DUDE/{target}") |
| np.save(f"{self.args.results_path}/DUDE/{target}/saved_mols_embed.npy", mol_reps) |
| np.save(f"{self.args.results_path}/DUDE/{target}/saved_target_embed.npy", pocket_reps) |
| np.save(f"{self.args.results_path}/DUDE/{target}/saved_labels.npy", labels) |
| auc, bedroc, ef_list, re_list = cal_metrics(labels, res_single, 80.5) |
|
|
| print(target) |
| print("ef:", ef_list) |
|
|
| return auc, bedroc, ef_list, re_list, res_single, labels |
|
|
| def test_dude_target_regression(self, target, model, seq, **kwargs): |
|
|
| data_path = f"{self.args.data}/DUD-E/" + target + "/mols_real.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = 64 |
| print(num_data // bsz) |
| mol_reps = [] |
| mol_names = [] |
| labels = [] |
|
|
| |
| print("begin with target:", target) |
| print("number of mol:", len(mol_dataset)) |
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, collate_fn=mol_dataset.collater) |
|
|
| act_preds_all = [] |
|
|
| |
|
|
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, collate_fn=mol_dataset.collater, |
| num_workers=8) |
| for _, mol_sample in enumerate(tqdm(mol_data)): |
| mol_sample = unicore.utils.move_to_cuda(mol_sample) |
| mol_names.extend(mol_sample["smi_name"]) |
| labels.extend(mol_sample["target"].detach().cpu().numpy()) |
|
|
| |
| data_path = f"{self.args.data}/DUD-E/" + target + "/pocket.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, |
| collate_fn=pocket_dataset.collater) |
| act_preds = [] |
| pocket_names = [] |
|
|
| for _, pocket_sample in enumerate(pocket_data): |
| pocket_sample = unicore.utils.move_to_cuda(pocket_sample) |
| pred = model.forward(protein_sequences=seq, **pocket_sample["net_input"], **mol_sample["net_input"]) |
| pocket_name = pocket_sample["pocket_name"] |
| act_preds.append(pred.detach().cpu().numpy()) |
| pocket_names.append(pocket_name) |
|
|
| act_preds = np.concatenate(act_preds, axis=0) |
| act_preds_all.append(act_preds) |
|
|
| res = np.concatenate(act_preds_all, axis=1) |
| res_single = res.max(axis=0) |
| os.system(f"mkdir -p {self.args.results_path}/DUDE/{target}") |
| np.save(f"{self.args.results_path}/DUDE/{target}/saved_labels.npy", labels) |
| np.save(f"{self.args.results_path}/DUDE/{target}/saved_preds.npy", res_single) |
| auc, bedroc, ef_list, re_list = cal_metrics(labels, res_single, 80.5) |
|
|
| print(target) |
| print("ef:", ef_list) |
|
|
| return auc, bedroc, ef_list, re_list, res_single, labels |
|
|
| def test_dude(self, model, **kwargs): |
|
|
| targets = list(os.listdir(f"{self.args.data}/DUD-E/")) |
| auc_list = [] |
| bedroc_list = [] |
| ef_list = [] |
| res_list = [] |
| labels_list = [] |
| re_list = { |
| "0.005": [], |
| "0.01": [], |
| "0.02": [], |
| "0.05": [], |
| } |
| ef_list = { |
| "0.005": [], |
| "0.01": [], |
| "0.02": [], |
| "0.05": [], |
| } |
| targets.reverse() |
| uniprot_list = json.load(open(f"{self.args.data}/dude.json")) |
| target2uniport = {x[2]: x[0] for x in uniprot_list} |
|
|
| for i, target in enumerate(targets): |
| seq = get_uniprot_seq(target2uniport[target.upper()]) |
| if self.args.arch in ["DTA", "pocketregression"]: |
| auc, bedroc, ef, re, res_single, labels = self.test_dude_target_regression(target, model, seq) |
| else: |
| auc, bedroc, ef, re, res_single, labels = self.test_dude_target(target, model, seq) |
| auc_list.append(auc) |
| bedroc_list.append(bedroc) |
| for key in ef: |
| ef_list[key].append(ef[key]) |
| for key in re_list: |
| re_list[key].append(re[key]) |
|
|
| print("auc mean", np.mean(auc_list)) |
| print("bedroc mean", np.mean(bedroc_list)) |
|
|
| for key in ef_list: |
| print("ef", key, "mean", np.mean(ef_list[key])) |
|
|
| for key in re_list: |
| print("re", key, "mean", np.mean(re_list[key])) |
|
|
| return |
|
|
| def test_dekois_target(self, target, model, seq, **kwargs): |
|
|
| data_path = f"{self.args.data}/DEKOIS_2.0x/{target}/{target}_lig.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = 64 |
| print(num_data // bsz) |
| mol_reps = [] |
| mol_names = [] |
| labels = [] |
|
|
| |
| print("begin with target:", target) |
| print("number of mol:", len(mol_dataset)) |
| mol_data = torch.utils.data.DataLoader(mol_dataset, num_workers=4, batch_size=bsz, |
| collate_fn=mol_dataset.collater) |
|
|
| for _, sample in enumerate(tqdm(mol_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| mol_emb = model.mol_forward(**sample["net_input"]) |
| mol_emb = mol_emb.detach().cpu().numpy() |
| |
| mol_reps.append(mol_emb) |
| mol_names.extend(sample["smi_name"]) |
| labels.extend(sample["target"].detach().cpu().numpy()) |
| mol_reps = np.concatenate(mol_reps, axis=0) |
| labels = np.array(labels, dtype=np.int32) |
| |
| data_path = f"{self.args.data}/DEKOIS_2.0x/{target}/{target}_pocket.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, collate_fn=pocket_dataset.collater) |
| pocket_reps = [] |
|
|
| for _, sample in enumerate(tqdm(pocket_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| pocket_emb = model.pocket_forward(protein_sequences=seq, **sample["net_input"]) |
| pocket_emb = pocket_emb.detach().cpu().numpy() |
| pocket_reps.append(pocket_emb) |
| pocket_reps = np.concatenate(pocket_reps, axis=0) |
| print(pocket_reps.shape) |
| res = pocket_reps @ mol_reps.T |
|
|
| res_single = res.max(axis=0) |
| os.system(f"mkdir -p {self.args.results_path}/DEKOIS/{target}") |
| print(f"writing to {self.args.results_path}/DEKOIS/{target}") |
| np.save(f"{self.args.results_path}/DEKOIS/{target}/saved_mols_embed.npy", mol_reps) |
| np.save(f"{self.args.results_path}/DEKOIS/{target}/saved_target_embed.npy", pocket_reps) |
| np.save(f"{self.args.results_path}/DEKOIS/{target}/saved_labels.npy", labels) |
| auc, bedroc, ef_list, re_list = cal_metrics(labels, res_single, 80.5) |
|
|
| print(target) |
| print("ef:", ef_list) |
|
|
| return auc, bedroc, ef_list, re_list, res_single, labels |
|
|
| def test_dekois_target_regression(self, target, model, seq, **kwargs): |
|
|
| data_path = f"{self.args.data}/DEKOIS_2.0x/{target}/{target}_lig.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = 64 |
| print(num_data // bsz) |
| mol_reps = [] |
| mol_names = [] |
| labels = [] |
| act_preds_all = [] |
|
|
| |
| print("begin with target:", target) |
| print("number of mol:", len(mol_dataset)) |
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, collate_fn=mol_dataset.collater, |
| num_workers=8) |
| for _, mol_sample in enumerate(tqdm(mol_data)): |
| mol_sample = unicore.utils.move_to_cuda(mol_sample) |
| mol_names.extend(mol_sample["smi_name"]) |
| labels.extend(mol_sample["target"].detach().cpu().numpy()) |
|
|
| |
| data_path = f"{self.args.data}/DEKOIS_2.0x/{target}/{target}_pocket.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, |
| collate_fn=pocket_dataset.collater) |
| act_preds = [] |
| pocket_names = [] |
|
|
| for _, pocket_sample in enumerate(pocket_data): |
| pocket_sample = unicore.utils.move_to_cuda(pocket_sample) |
| pred = model.forward(protein_sequences=seq, **pocket_sample["net_input"], **mol_sample["net_input"]) |
| pocket_name = pocket_sample["pocket_name"] |
| act_preds.append(pred.detach().cpu().numpy()) |
| pocket_names.append(pocket_name) |
|
|
| act_preds = np.concatenate(act_preds, axis=0) |
| act_preds_all.append(act_preds) |
|
|
| res = np.concatenate(act_preds_all, axis=1) |
|
|
| res_single = res.max(axis=0) |
| os.system(f"mkdir -p {self.args.results_path}/DEKOIS/{target}") |
| print(f"writing to {self.args.results_path}/DEKOIS/{target}") |
| np.save(f"{self.args.results_path}/DEKOIS/{target}/saved_labels.npy", labels) |
| np.save(f"{self.args.results_path}/DEKOIS/{target}/saved_preds.npy", res_single) |
| auc, bedroc, ef_list, re_list = cal_metrics(labels, res_single, 80.5) |
|
|
| print(target) |
| print("ef:", ef_list) |
|
|
| return auc, bedroc, ef_list, re_list, res_single, labels |
|
|
| def test_dekois(self, model, **kwargs): |
|
|
| targets = list(os.listdir(f"{self.args.data}/DEKOIS_2.0x/")) |
| auc_list = [] |
| bedroc_list = [] |
| ef_list = [] |
| res_list = [] |
| labels_list = [] |
| re_list = { |
| "0.005": [], |
| "0.01": [], |
| "0.02": [], |
| "0.05": [], |
| } |
| ef_list = { |
| "0.005": [], |
| "0.01": [], |
| "0.02": [], |
| "0.05": [], |
| } |
| targets.reverse() |
|
|
| uniprot_list = json.load(open(f"{self.args.data}/dekois.json")) |
| target2uniport = {x[2]: x[0] for x in uniprot_list} |
|
|
| for i, target in enumerate(targets): |
| if not os.path.exists(f"{self.args.data}/DEKOIS_2.0x/{target}/{target}_lig.lmdb"): |
| continue |
| seq = get_uniprot_seq(target2uniport[target.upper()]) |
| if self.args.arch in ["DTA", "pocketregression"]: |
| auc, bedroc, ef, re, res_single, labels = self.test_dekois_target_regression(target, model, seq) |
| else: |
| auc, bedroc, ef, re, res_single, labels = self.test_dekois_target(target, model, seq) |
| auc_list.append(auc) |
| bedroc_list.append(bedroc) |
| for key in ef: |
| ef_list[key].append(ef[key]) |
| for key in re_list: |
| re_list[key].append(re[key]) |
| |
| |
| |
|
|
| print("auc mean", np.mean(auc_list)) |
| print("bedroc mean", np.mean(bedroc_list)) |
|
|
| for key in ef_list: |
| print("ef", key, "mean", np.mean(ef_list[key])) |
|
|
| for key in re_list: |
| print("re", key, "mean", np.mean(re_list[key])) |
|
|
| return |
|
|
| def test_demo(self, model): |
| data_path = self.args.demo_lig_file |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
|
|
| bsz = self.args.batch_size |
| mol_reps = [] |
| mol_smis = [] |
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, |
| collate_fn=mol_dataset.collater, |
| num_workers=self.args.num_workers) |
|
|
| for _, sample in enumerate(mol_data): |
| sample = unicore.utils.move_to_cuda(sample) |
| mol_emb = model.mol_forward(**sample["net_input"]) |
| mol_emb = mol_emb.detach().cpu().numpy() |
| |
| mol_reps.append(mol_emb) |
| mol_smis.extend(sample["smi_name"]) |
| mol_reps = np.concatenate(mol_reps, axis=0) |
|
|
| data_path = self.args.demo_prot_file |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, collate_fn=pocket_dataset.collater) |
| sample = list(pocket_data)[0] |
|
|
| sample = unicore.utils.move_to_cuda(sample) |
| seq = get_uniprot_seq(self.args.demo_uniprot) |
| pocket_emb = model.pocket_forward(protein_sequences=seq, **sample["net_input"]) |
| pocket_reps = pocket_emb.detach().cpu().numpy() |
|
|
| os.system(f"mkdir -p {self.args.results_path}") |
| json.dump(mol_smis, open(f"{self.args.results_path}/saved_smis.json", "w")) |
| np.save(f"{self.args.results_path}/saved_mols_embed.npy", mol_reps) |
| np.save(f"{self.args.results_path}/saved_target_embed.npy", pocket_reps) |
|
|
| def test_fep_target(self, target, model, label_info, **kwargs): |
| data_path = f"{self.args.data}/FEP/lmdbs/{target}_lig.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = 64 |
| mol_reps = [] |
| mol_smis = [] |
| labels = [] |
|
|
| |
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, collate_fn=mol_dataset.collater) |
|
|
| for _, sample in enumerate(mol_data): |
| sample = unicore.utils.move_to_cuda(sample) |
| mol_emb = model.mol_forward(**sample["net_input"]) |
| mol_emb = mol_emb.detach().cpu().numpy() |
| |
| mol_reps.append(mol_emb) |
| mol_smis.extend(sample["smi_name"]) |
| mol_reps = np.concatenate(mol_reps, axis=0) |
| |
| data_path = f"{self.args.data}/FEP/lmdbs/{target}.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, collate_fn=pocket_dataset.collater) |
| pocket_reps = [] |
|
|
| for _, sample in enumerate(pocket_data): |
| sample = unicore.utils.move_to_cuda(sample) |
| seq = label_info["sequence"] |
| pocket_emb = model.pocket_forward(protein_sequences=seq, **sample["net_input"]) |
| pocket_emb = pocket_emb.detach().cpu().numpy() |
| pocket_reps.append(pocket_emb) |
|
|
| pocket_reps = np.concatenate(pocket_reps, axis=0) |
| res = pocket_reps @ mol_reps.T |
|
|
| res_single = res.max(axis=0) |
| act_dict = {} |
| for lig in label_info["ligands"]: |
| act_dict[lig["smi"]] = float(lig["act"]) |
| real_dg = np.array([act_dict[smi] for smi in mol_smis]) |
| pred_dg = res_single |
| from scipy import stats |
| corr = stats.pearsonr(real_dg, pred_dg).statistic |
| if corr < 0: |
| r2 = 0 |
| else: |
| r2 = corr ** 2 |
|
|
| os.system(f"mkdir -p {self.args.results_path}/FEP/{target}") |
| np.save(f"{self.args.results_path}/FEP/{target}/saved_mols_embed.npy", mol_reps) |
| np.save(f"{self.args.results_path}/FEP/{target}/saved_target_embed.npy", pocket_reps) |
| np.save(f"{self.args.results_path}/FEP/{target}/saved_labels.npy", real_dg) |
| json.dump(mol_smis, open(f"{self.args.results_path}/FEP/{target}/saved_smis.json", "w")) |
| return r2 |
|
|
| def test_fep_target_regression(self, target, model, label_info, **kwargs): |
|
|
| data_path = f"{self.args.data}/FEP/lmdbs/{target}_lig.lmdb" |
| mol_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(mol_dataset) |
| bsz = 64 |
| act_preds_all = [] |
| mol_smis = [] |
| labels = [] |
|
|
| |
| mol_data = torch.utils.data.DataLoader(mol_dataset, batch_size=bsz, collate_fn=mol_dataset.collater) |
| for _, mol_sample in enumerate(tqdm(mol_data)): |
| mol_sample = unicore.utils.move_to_cuda(mol_sample) |
| mol_smis.extend(mol_sample["smi_name"]) |
| labels.extend(mol_sample["target"].detach().cpu().numpy()) |
|
|
| |
| data_path = f"{self.args.data}/FEP/lmdbs/{target}.lmdb" |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, batch_size=bsz, |
| collate_fn=pocket_dataset.collater) |
| act_preds = [] |
| pocket_names = [] |
|
|
| for _, pocket_sample in enumerate(tqdm(pocket_data)): |
| pocket_sample = unicore.utils.move_to_cuda(pocket_sample) |
| seq = label_info["sequence"] |
| pred = model.forward(protein_sequences=seq, **pocket_sample["net_input"], **mol_sample["net_input"]) |
| pocket_name = pocket_sample["pocket_name"] |
| pocket_names.append(pocket_name) |
| print("pred", pred.shape) |
| act_preds.append(pred.detach().cpu().numpy() + 6.) |
|
|
| act_preds = np.concatenate(act_preds, axis=0) |
| print("act_preds", act_preds.shape) |
| act_preds_all.append(act_preds) |
|
|
| res = np.concatenate(act_preds_all, axis=1) |
| res_single = res.max(axis=0) |
| act_dict = {} |
| for lig in label_info["ligands"]: |
| act_dict[lig["smi"]] = float(lig["act"]) |
| real_dg = np.array([act_dict[smi] for smi in mol_smis]) |
| pred_dg = res_single |
| from scipy import stats |
| corr = stats.pearsonr(real_dg, pred_dg).statistic |
| |
| |
| |
| |
|
|
| os.system(f"mkdir -p {self.args.results_path}/FEP/{target}") |
| np.save(f"{self.args.results_path}/FEP/{target}/saved_labels.npy", real_dg) |
| np.save(f"{self.args.results_path}/FEP/{target}/saved_preds.npy", pred_dg) |
| json.dump(mol_smis, open(f"{self.args.results_path}/FEP/{target}/saved_smis.json", "w")) |
|
|
| return corr |
|
|
| def test_fep(self, model, **kwargs): |
| labels_fep = json.load( |
| open(f"{self.args.data}/FEP/fep_labels.json")) |
| ligands_dict = {x["pockets"][0]: x for x in labels_fep} |
| rho_list = [] |
| for i, target in enumerate(ligands_dict.keys()): |
| if self.args.arch in ["DTA", "pocketregression"]: |
| rho = self.test_fep_target_regression(target, model, ligands_dict[target]) |
| else: |
| rho = self.test_fep_target(target, model, ligands_dict[target]) |
| |
| rho_list.append(rho) |
|
|
| print(self.args.results_path.split("/")[-1], np.mean(rho_list), np.median(rho_list)) |
|
|
| def inference_pdbbind(self, model, split="train", **kwargs): |
| pdbbind_dataset = self.load_dataset(split, load_name=True, shuffle=False) |
| num_data = len(pdbbind_dataset) |
| bsz = 32 |
| print(num_data // bsz) |
| mol_reps = [] |
| pocket_reps = [] |
| pdbbind_ids = [] |
| mol_smis = [] |
| pocket2seq = {} |
| if split == "train": |
| pdbbind_label = json.load(open(f"{self.args.data}/train_label_pdbbind_seq.json")) |
| else: |
| pdbbind_label = json.load(open(f"/casf_label_seq.json")) |
| for assay in pdbbind_label: |
| seq = assay["sequence"] |
| pockets = assay["pockets"] |
| for pocket in pockets: |
| pocket2seq[pocket.split("_")[0]] = seq |
|
|
| |
| print("number of data:", len(pdbbind_dataset)) |
| mol_data = torch.utils.data.DataLoader(pdbbind_dataset, num_workers=8, batch_size=bsz, |
| collate_fn=pdbbind_dataset.collater) |
|
|
| for _, sample in enumerate(tqdm(mol_data)): |
| |
| sample = unicore.utils.move_to_cuda(sample) |
| pocket_names = sample["pocket_name"] |
| seq = [pocket2seq[x] for x in pocket_names] |
| mol_emb, pocket_emb, _, _ = model.forward(**sample["net_input"], protein_sequences=seq) |
| mol_emb = mol_emb[0].detach().cpu().numpy() |
| mol_reps.append(mol_emb) |
| pocket_emb = pocket_emb[0].detach().cpu().numpy() |
| pocket_reps.append(pocket_emb) |
|
|
| pdbbind_ids += sample["pocket_name"] |
| mol_smis += sample["smi_name"] |
|
|
| mol_reps = np.concatenate(mol_reps, axis=0) |
| pocket_reps = np.concatenate(pocket_reps, axis=0) |
|
|
| write_dir = f"{self.args.results_path}/PDBBind" |
| if not os.path.exists(write_dir): |
| os.system(f"mkdir -p {write_dir}") |
| np.save(f"{write_dir}/{split}_mol_reps.npy", mol_reps) |
| np.save(f"{write_dir}/{split}_pocket_reps.npy", pocket_reps) |
| json.dump(pdbbind_ids, open(f"{write_dir}/{split}_pdbbind_ids.json", "w")) |
| json.dump(mol_smis, open(f"{write_dir}/{split}_mol_smis.json", "w")) |
|
|
| def inference_bdb_lig(self, model): |
| data_path = f"{self.args.data}/train_lig_all_blend.lmdb" |
| bdb_dataset = self.load_mols_dataset(data_path, "atoms", "coordinates") |
| num_data = len(bdb_dataset) |
| bsz = 128 |
| print(num_data // bsz) |
| |
| print("number of data:", len(bdb_dataset)) |
|
|
| mol_reps = [] |
| mol_smis = [] |
| mol_data = torch.utils.data.DataLoader(bdb_dataset, num_workers=8, batch_size=bsz, |
| collate_fn=bdb_dataset.collater) |
| for _, sample in enumerate(tqdm(mol_data)): |
| |
| sample = unicore.utils.move_to_cuda(sample) |
| mol_emb = model.mol_forward(**sample["net_input"]) |
| mol_emb = mol_emb.detach().cpu().numpy() |
| mol_reps.append(mol_emb) |
| mol_smis += sample["smi_name"] |
|
|
| mol_reps = np.concatenate(mol_reps, axis=0) |
|
|
| write_dir = f"{self.args.results_path}/BDB" |
| if not os.path.exists(write_dir): |
| os.mkdir(write_dir) |
| np.save(f"{write_dir}/bdb_mol_reps.npy", mol_reps) |
| json.dump(mol_smis, open(f"{write_dir}/bdb_mol_smis.json", "w")) |
|
|
| def inference_bdb_pocket(self, model): |
| data_path = f"{self.args.data}/train_prot_all_blend.lmdb" |
| blend_label = json.load(open(f"{self.args.data}/train_label_blend_seq_full.json")) |
| pocket_dataset = self.load_pockets_dataset(data_path) |
| bsz = 32 |
| pocket_data = torch.utils.data.DataLoader(pocket_dataset, num_workers=8, batch_size=bsz, |
| collate_fn=pocket_dataset.collater) |
| pocket_reps = [] |
| pocket_names = [] |
| pocket2seq = {} |
| for assay in blend_label: |
| seq = assay["sequence"] |
| pockets = assay["pockets"] |
| for pocket in pockets: |
| pocket2seq[pocket] = seq |
|
|
| for _, sample in enumerate(tqdm(pocket_data)): |
| sample = unicore.utils.move_to_cuda(sample) |
| pocket_name = sample["pocket_name"] |
| seq_list = [pocket2seq.get(x, "") for x in pocket_name] |
|
|
| pocket_emb = model.pocket_forward(protein_sequences=seq_list, **sample["net_input"]) |
| pocket_emb = pocket_emb.detach().cpu().numpy() |
| for seq, emb, name in zip(seq_list, pocket_emb, sample["pocket_name"]): |
| if seq != "": |
| pocket_names.append(name) |
| pocket_reps.append(emb) |
| pocket_reps = np.stack(pocket_reps, axis=0) |
| print(pocket_reps.shape) |
| write_dir = f"{self.args.results_path}/BDB" |
| if not os.path.exists(write_dir): |
| os.mkdir(write_dir) |
| np.save(f"{write_dir}/bdb_pocket_reps.npy", pocket_reps) |
| json.dump(pocket_names, open(f"{write_dir}/bdb_pocket_names.json", "w")) |