File size: 8,140 Bytes
8efb4bd | 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 | import pickle
from functools import lru_cache
import os
from typing import Dict, List
import numpy as np
from automatic_pipeline.configurable import get_scores_from_path
from automatic_pipeline.libs.get_alphafold_jobs import get_residue_diff
from automatic_pipeline.libs.utils_classes import PdbPath, AlphaFoldJobInfo, AFResultScoredPair, SubunitsInfo, \
SubunitPdbInfo, AFSubunitScores, AFInteractionScores
from automatic_pipeline.libs.utils_pdb import get_pdb_model_readonly, are_subunits_close_in_pdb, get_interface_res_ids
try:
import ujson as json
except ModuleNotFoundError:
print("using slow json!")
import json
@lru_cache(5)
def get_alphafold_scores(pdb_path: str):
return json.load(open(get_scores_from_path(pdb_path), "r"))
def parse_results_to_scored_pairs(alphafold_results: Dict[AlphaFoldJobInfo, List[PdbPath]],
subunits_info: SubunitsInfo) -> List[AFResultScoredPair]:
all_pairs = []
for job_count, (af_job_info, pdb_paths) in enumerate(alphafold_results.items()):
if job_count % 5 == 0:
print("finding pairs", job_count, "out of", len(alphafold_results), "found", len(all_pairs))
for pdb_path in pdb_paths:
pdb_model = get_pdb_model_readonly(pdb_path)
chain_names = sorted([c.id for c in pdb_model.get_chains()])
subunits_per_chain = [[af_job_info.subunit_names[0]]]
for subunit_name, should_merge in zip(af_job_info.subunit_names[1:], af_job_info.merged_subunits):
if should_merge:
subunits_per_chain[-1].append(subunit_name)
else:
subunits_per_chain.append([subunit_name])
assert len(chain_names) == len(subunits_per_chain), f"Mismatch subunits and chains lengths " \
f"{len(chain_names)} {len(subunits_per_chain)}"
subunits_pdb_infos = []
pdb_start_res = 0
for chain_name, subunit_list in zip(chain_names, subunits_per_chain):
chain_start_res = 1
for subunit_i in range(len(subunit_list)):
subunit_name = subunit_list[subunit_i]
subunit = subunits_info[subunit_name]
if subunit_i > 0:
prev_subunit = subunits_info[subunit_list[subunit_i - 1]]
chain_start_res += get_residue_diff(prev_subunit, subunit)
subunits_pdb_infos.append(SubunitPdbInfo(chain_id=chain_name, chain_residue_id=chain_start_res,
pdb_residue_id=pdb_start_res,
length=len(subunit.sequence)))
chain_start_res += len(subunit.sequence)
pdb_start_res += len(subunit.sequence)
for i in range(len(subunits_pdb_infos)):
for j in range(i + 1, len(subunits_pdb_infos)):
if are_subunits_close_in_pdb(pdb_path, subunits_pdb_infos[i], subunits_pdb_infos[j]):
all_pairs.append(AFResultScoredPair(pdb_path=pdb_path,
subunits_names=(af_job_info.subunit_names[i],
af_job_info.subunit_names[j]),
subunit1_pdb_info=subunits_pdb_infos[i],
subunit2_pdb_info=subunits_pdb_infos[j]))
return all_pairs
def _get_subunit_scores(plddt: np.ndarray, pae: np.ndarray, subunit_indexes: List[int], interface_indexes: List[int]) \
-> AFSubunitScores:
subunit_plddt = plddt[subunit_indexes]
return AFSubunitScores(
plddt_avg=float(np.mean(subunit_plddt)),
plddt_percentile=list(np.percentile(subunit_plddt, np.arange(0, 101, 10))),
plddt_interface_avg=float(np.mean(plddt[interface_indexes])),
plddt_interface_percentile=list(np.percentile(plddt[interface_indexes], np.arange(0, 101, 10))),
self_pae_avg=float(np.mean(pae[subunit_indexes][:, subunit_indexes])),
self_pae_percentile=list(np.percentile(pae[subunit_indexes][:, subunit_indexes], np.arange(0, 101, 10)))
)
def _get_interaction_scores(pae: np.ndarray, subunit_indexes1: List[int], subunit_indexes2: List[int],
interface_indexes1: List[int], interface_indexes2: List[int]) -> AFInteractionScores:
all_pae = np.concatenate([pae[subunit_indexes1][:, subunit_indexes2].flatten(),
pae[subunit_indexes2][:, subunit_indexes1].flatten()])
pae_joined_interface = np.concatenate([pae[interface_indexes1][:, interface_indexes2].flatten(),
pae[interface_indexes2][:, interface_indexes1].flatten()])
return AFInteractionScores(
pae_avg=float(np.mean(all_pae)),
pae_percentile=list(np.percentile(all_pae, np.arange(0, 101, 10))),
pae_joined_interface_avg=float(np.mean(pae_joined_interface)),
pae_joined_interface_percentile=list(np.percentile(pae_joined_interface, np.arange(0, 101, 10))),
interface1_size=len(interface_indexes1),
interface2_size=len(interface_indexes2),
)
def score_af_results_as_pairs(alphafold_results: Dict[AlphaFoldJobInfo, List[PdbPath]],
output_path: str, subunits_info: SubunitsInfo) -> List[AFResultScoredPair]:
if os.path.exists(output_path):
with open(output_path, "rb") as f:
all_pairs = pickle.load(f)
else:
all_pairs = parse_results_to_scored_pairs(alphafold_results, subunits_info)
pickle.dump(all_pairs, open(output_path, "wb"))
not_scored_results: List[AFResultScoredPair] = [i for i in all_pairs if i.subunit1_scores is None]
# enrich with scores
for count, pair in enumerate(not_scored_results):
# pdb_info = AlphaFoldJobInfo.from_result_path(pair.pdb_path)
job_scores = get_alphafold_scores(pair.pdb_path)
plddt = np.array(job_scores["plddt"])
pae = np.array(job_scores["pae"])
# index are 0-based
subunit1_indexes = [i for i in range(pair.subunit1_pdb_info.length)
if subunits_info[pair.subunits_names[0]].sequence[i] != "X"]
subunit1_indexes = np.array(subunit1_indexes) + pair.subunit1_pdb_info.pdb_residue_id
subunit2_indexes = [i for i in range(pair.subunit2_pdb_info.length)
if subunits_info[pair.subunits_names[1]].sequence[i] != "X"]
subunit2_indexes = np.array(subunit2_indexes) + pair.subunit2_pdb_info.pdb_residue_id
subunit1_interface, subunit2_interface = get_interface_res_ids(pair.pdb_path,
pair.subunit1_pdb_info,
pair.subunit2_pdb_info)
subunit1_interface_indexes = np.array(list(subunit1_interface)) + pair.subunit1_pdb_info.pdb_residue_id
subunit2_interface_indexes = np.array(list(subunit2_interface)) + pair.subunit2_pdb_info.pdb_residue_id
pair.subunit1_scores = _get_subunit_scores(plddt, pae, subunit1_indexes, subunit1_interface_indexes)
pair.subunit2_scores = _get_subunit_scores(plddt, pae, subunit2_indexes, subunit2_interface_indexes)
pair.interaction_scores = _get_interaction_scores(pae, subunit1_indexes, subunit2_indexes,
subunit1_interface_indexes, subunit2_interface_indexes)
if count % 100 == 0 or count == len(not_scored_results) - 1:
print(f"Processed {count + 1}/{len(not_scored_results)} total pairs: {len(all_pairs)}")
pickle.dump(all_pairs, open(output_path, "wb"))
return all_pairs
|