File size: 9,795 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
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
import os
import shutil
from typing import List, Dict, Tuple

from automatic_pipeline.configurable import AF2TRANS_BIN_PATH, is_bash_running, COMB_ASSEMBLY_BIN_PATH, run_bash_file
from automatic_pipeline.libs.utils_classes import AFResultScoredPair, SubunitsInfo, SubunitName
from automatic_pipeline.libs.utils_pdb import extract_pdb_info, copy_pdb_rename_chain, copy_pdb_set_start_offset


def extract_ref_structs(af_scored_pairs: List[AFResultScoredPair], output_path: str, subunits_info: SubunitsInfo):
    all_subunit_names = sum([i.get_chained_names() for i in subunits_info.values()], [])
    if all([os.path.exists(os.path.join(output_path, f"{i}.pdb")) for i in all_subunit_names]):
        return
    os.makedirs(output_path, exist_ok=True)

    ref_structs: Dict[SubunitName, Tuple[float, AFResultScoredPair, int]] = {}
    for result in af_scored_pairs:
        score = result.subunit1_scores.plddt_avg
        if result.subunits_names[0] not in ref_structs or ref_structs[result.subunits_names[0]][0] < score:
            ref_structs[result.subunits_names[0]] = (score, result, 0)

        score = result.subunit2_scores.plddt_avg
        if result.subunits_names[1] not in ref_structs or ref_structs[result.subunits_names[1]][0] < score:
            ref_structs[result.subunits_names[1]] = (score, result, 1)

    assert all(i in ref_structs for i in subunits_info.keys()), f"Missing ref structs {list(ref_structs.keys())} " \
                                                                f"{list(subunits_info.keys())}"
    for subunit_name, (score, result, subunit_idx) in ref_structs.items():
        subunit_info = subunits_info[subunit_name]
        subunit_pdb_info = result.subunit1_pdb_info if subunit_idx == 0 else result.subunit2_pdb_info
        print(f"ref_struct {subunit_name} {score} {subunit_pdb_info}")

        ref_struct_path = os.path.join(output_path, f"{subunit_name}.pdb")
        extract_pdb_info(result.pdb_path, subunit_pdb_info, ref_struct_path)
        copy_pdb_set_start_offset(ref_struct_path, subunit_info.start_res - 1, ref_struct_path)
        for chain_name, ident_subunit_name in zip(subunit_info.chain_names, subunit_info.get_chained_names()):
            copy_pdb_rename_chain(ref_struct_path, chain_name,
                                  os.path.join(output_path, f"{ident_subunit_name}.pdb"))
        os.remove(ref_struct_path)


def _score_result_transform(result: AFResultScoredPair) -> float:
    subunit1_score = result.subunit1_scores
    subunit2_score = result.subunit2_scores
    interaction_score = result.interaction_scores
    full_pae_avg: float = (interaction_score.pae_avg + subunit1_score.self_pae_avg + subunit2_score.self_pae_avg) / 3

    # when full_pae_avg=20, we will return 1. has quadric properties
    return max([1, 100 - (1/4) * (full_pae_avg ** 2)])


def create_transformations(af_scored_pairs: List[AFResultScoredPair], ref_structs_folder: str,

                           subunits_info: SubunitsInfo, transformations_output_folder: str):
    os.makedirs(transformations_output_folder, exist_ok=True)
    os.chdir(transformations_output_folder)

    temp_dirname = "temp_pdbs_dir"
    temp_dir = os.path.join(transformations_output_folder, temp_dirname)

    total_pairs_used = 0
    ordered_subunits = sorted(list(subunits_info.keys()))
    for i in range(len(ordered_subunits)):
        for j in range(i, len(ordered_subunits)):
            os.makedirs(temp_dir, exist_ok=True)
            subunit_name1 = ordered_subunits[i]
            subunit_name2 = ordered_subunits[j]
            chained_subunit_name1 = subunits_info[subunit_name1].get_chained_names()[0]
            chained_subunit_name2 = subunits_info[subunit_name2].get_chained_names()[0]

            if i == j and len(subunits_info[subunit_name1].chain_names) == 1:
                continue

            output_file_path = os.path.join(transformations_output_folder,
                                            f"{subunit_name1}_plus_{subunit_name2}")
            if all([os.path.exists(os.path.join(transformations_output_folder,
                                                f"{chained_subunit_name1}_plus_{chained_subunit_name2}"))
                    for chained_subunit_name1 in subunits_info[subunit_name1].get_chained_names()
                    for chained_subunit_name2 in subunits_info[subunit_name2].get_chained_names()]):
                continue

            results_for_pair = [result for result in af_scored_pairs
                                if result.subunits_names == (subunit_name1, subunit_name2)]
            if len(results_for_pair) == 0:
                print("No results for pair", subunit_name1, subunit_name2)
                continue
            total_pairs_used += len(results_for_pair)

            cmd = AF2TRANS_BIN_PATH + " "
            ref_subunit1_path = os.path.join(ref_structs_folder, chained_subunit_name1 + ".pdb")
            ref_subunit2_path = os.path.join(ref_structs_folder, chained_subunit_name2 + ".pdb")
            cmd += f"{ref_subunit1_path} {ref_subunit2_path} "

            results_for_pair.sort(key=_score_result_transform, reverse=True)
            scores = []
            for counter, result in enumerate(results_for_pair):
                chain1_pdb = os.path.join(temp_dir, f"{counter}_0_{os.path.basename(result.pdb_path)}")
                chain2_pdb = os.path.join(temp_dir, f"{counter}_1_{os.path.basename(result.pdb_path)}")
                extract_pdb_info(result.pdb_path, result.subunit1_pdb_info, chain1_pdb)
                extract_pdb_info(result.pdb_path, result.subunit2_pdb_info, chain2_pdb)

                cmd += f"{chain1_pdb} {chain2_pdb} "
                scores.append(_score_result_transform(result))
            print(f"result count for {subunit_name1} {subunit_name2} {len(results_for_pair)}")
            os.system(f"{cmd} > {output_file_path}")
            print(f"Saved transformations to {output_file_path}")

            # Calculate and set score based on PAE for transformations
            content_lines = open(output_file_path, "r").read().split("\n")
            scored_lines = []
            assert len(content_lines) == len(scores) + 1
            for line, score in zip(content_lines, scores):
                if len(line) == 0:
                    continue
                splitted = line.split(" | ")
                scored_lines.append(" | ".join([splitted[0], str(score)] + splitted[2:]))
            open(output_file_path, "w").write("\n".join(scored_lines) + "\n")

            # copy results for ident chains
            aliases_c1 = subunits_info[subunit_name1].get_chained_names()
            aliases_c2 = subunits_info[subunit_name2].get_chained_names()

            for c1 in range(len(aliases_c1)):
                start_from = c1 + 1 if subunit_name1 == subunit_name2 else 0
                for c2 in range(start_from, len(aliases_c2)):
                    alias_output_path = os.path.join(transformations_output_folder,
                                                     f"{aliases_c1[c1]}_plus_{aliases_c2[c2]}")
                    shutil.copy(output_file_path, alias_output_path)
            os.remove(output_file_path)
            shutil.rmtree(temp_dir)
    print("used a total of", total_pairs_used, "pairs out of ", len(af_scored_pairs))


def run_assembly(af_scored_pairs: List[AFResultScoredPair], subunits_info: SubunitsInfo, output_path: str) -> str:
    assembly_output_path = os.path.join(output_path, "assembly_output")
    transformations_folder = os.path.join(output_path, "transformations")

    os.makedirs(assembly_output_path, exist_ok=True)
    os.makedirs(transformations_folder, exist_ok=True)
    run_file_path = os.path.join(assembly_output_path, "run.sh")

    done_file_path = os.path.join(assembly_output_path, "assembly_done")
    if os.path.exists(done_file_path):
        return "success"

    if is_bash_running(run_file_path):
        return "running"

    # extract ref structs
    extract_ref_structs(af_scored_pairs, assembly_output_path, subunits_info)

    # extract transformations
    create_transformations(af_scored_pairs, assembly_output_path, subunits_info, transformations_folder)

    # prepare hasp input files
    with open(os.path.join(assembly_output_path, "chain.list"), "w") as f:
        sorted_all_subunits = sorted(sum([i.get_chained_names() for i in subunits_info.values()], []))
        for chained_subunit_name in sorted_all_subunits:
            f.write(f"{chained_subunit_name}.pdb\n")
    os.chdir(assembly_output_path)
    open("xlink_consts.txt", "w").close()

    with open(run_file_path, "w") as f:
        f.write("#!/bin/bash\n")
        f.write(f"cd {assembly_output_path}\n")
        f.write(f"{COMB_ASSEMBLY_BIN_PATH} chain.list {transformations_folder}/ 900 100 xlink_consts.txt "
                f"-b 0.05 -t 80\n")
        f.write(f"touch {done_file_path}\n")
    run_bash_file(run_file_path)
    print("Started assembly")
    return "running"


def get_assembly_results(output_path: str) -> float:
    clusters_path = os.path.join(output_path, "assembly_output", "output_clustered.res")
    if not os.path.exists(clusters_path):
        return 0

    results_as_str = [i for i in open(clusters_path, "r").read().split("\n") if len(i) > 0]
    scores = []
    for i, result_as_str in enumerate(results_as_str):
        splitted_result = result_as_str.split(" ")
        # score = float(splitted_result[splitted_result.index("transScore_") + 1])
        scores.append(float(splitted_result[splitted_result.index("weightedTransScore") + 1]))
    return max(scores, default=0)