File size: 24,415 Bytes
84ff331
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
import dataclasses
import os
import shutil
import subprocess
import sys
from collections import defaultdict
from typing import Dict, List, Tuple, Optional
import numpy as np
import scipy.spatial.distance

import Bio.SeqUtils
import Bio.PDB, Bio.PDB.Residue

from libs.prepare_complex import create_complexes
from libs.utils_classes import read_subunits_info, SubunitName, SubunitsInfo, INTERFACE_MIN_ATOM_DIST
from libs.utils_pdb import get_pdb_model_readonly, copy_pdb_set_start_offset, copy_pdb_rename_chain


THIS_SCRIPT_PATH = os.path.abspath(__file__)
BASE_PATH = os.path.dirname(THIS_SCRIPT_PATH)
REPO_ROOT = os.path.abspath(os.path.join(BASE_PATH, ".."))

BINARY_PATH = os.path.join(REPO_ROOT, "model", "CombinatorialAssembler")
AF2TRANS_BIN_PATH = os.path.join(BINARY_PATH, "AF2trans.out")
COMB_ASSEMBLY_BIN_PATH = os.path.join(BINARY_PATH, "CombinatorialAssembler.out")


# In most cases PartialSubunit will be the complete subunit, but this allows to input PDBs of interactions with only the
# interfaces between subunits.
@dataclasses.dataclass
class PartialSubunit:
    subunit_name: str
    pdb_path: str
    chain_id: str
    start_residue_id: int  # inclusive
    end_residue_id: int  # inclusive
    subunit_start_sequence_id: int
    is_complete: bool = False


@dataclasses.dataclass
class TransformationInfo:
    subunit_names: Tuple[str, str]
    pdb_path: str
    pdb_chain_ids: Tuple[str, str]
    rep_imposed_rmsds: Tuple[float, float]
    transformation_numbers: str
    score: float


def get_chain_to_seq(pdb_path: str, use_seqres: bool = True) -> Dict[str, str]:
    if use_seqres:
        chain_to_seq = {str(record.id): str(record.seq) for record in Bio.SeqIO.parse(pdb_path, 'pdb-seqres')}
        if len(chain_to_seq) > 0:
            return chain_to_seq

    model = get_pdb_model_readonly(pdb_path)
    chain_to_seq = {}
    for chain in model.get_chains():
        res_id_to_res = {res.get_id()[1]: res for res in chain.get_residues() if "CA" in res}
        if len(res_id_to_res) == 0:
            print("skipping empty chain", chain.get_id())
            continue
        chain_to_seq[chain.get_id()] = ""
        for i in range(1, max(res_id_to_res) + 1):
            if i in res_id_to_res:
                chain_to_seq[chain.get_id()] += Bio.SeqUtils.seq1(res_id_to_res[i].get_resname())
            else:
                chain_to_seq[chain.get_id()] += "X"
    return chain_to_seq


def _get_partial_subunit_residues(partial_subunit: PartialSubunit) -> List[Bio.PDB.Residue.Residue]:
    pdb_model = get_pdb_model_readonly(partial_subunit.pdb_path)
    return [res for res in pdb_model[partial_subunit.chain_id] if
            partial_subunit.start_residue_id <= res.id[1] <= partial_subunit.end_residue_id]


def extract_partial_subunit(partial_subunit: PartialSubunit, output_path: str):
    pdb_parser = Bio.PDB.PDBParser(QUIET=True)
    pdb_struct = pdb_parser.get_structure("original_pdb", partial_subunit.pdb_path)
    assert len(list(pdb_struct)) == 1, "can't extract if more than one model"
    model = next(iter(pdb_struct))
    chains = list(model.get_chains())
    assert len([c for c in chains if c.id == partial_subunit.chain_id]) == 1, f"Missing: {partial_subunit.chain_id}"
    for chain in chains:
        if chain.id != partial_subunit.chain_id:
            model.detach_child(chain.id)

    res_to_keep = _get_partial_subunit_residues(partial_subunit)

    res_to_remove = [res for res in model.get_residues() if res not in res_to_keep]
    for res in res_to_remove:
        res.parent.detach_child(res.id)

    io = Bio.PDB.PDBIO()
    io.set_structure(pdb_struct)
    io.save(output_path)


def extract_partial_from_representative(partial_subunit: PartialSubunit, representative_subunits_path: str,

                                        output_folder: str, subunits_info: SubunitsInfo) -> str:
    subunit_info = subunits_info[partial_subunit.subunit_name]
    start_res_id = partial_subunit.subunit_start_sequence_id + subunit_info.start_res
    end_res_id = start_res_id + (partial_subunit.end_residue_id - partial_subunit.start_residue_id)

    output_pdb_path = os.path.join(output_folder, f"{partial_subunit.subunit_name}_{start_res_id}_"
                                                  f"{end_res_id}.pdb")
    if os.path.exists(output_pdb_path):
        return output_pdb_path

    chain_name, ident_subunit_name = subunit_info.chain_names[0], subunit_info.get_chained_names()[0]
    rep_subunit_path = os.path.join(representative_subunits_path, f"{ident_subunit_name}.pdb")

    rep_partial_subunit = PartialSubunit(subunit_name=subunit_info.name,
                                         pdb_path=rep_subunit_path,
                                         chain_id=chain_name,
                                         start_residue_id=start_res_id,
                                         end_residue_id=end_res_id,
                                         subunit_start_sequence_id=0)  # subunit_start_sequence_id is ignored
    extract_partial_subunit(rep_partial_subunit, output_pdb_path)
    return output_pdb_path


def score_transformation(pdb_path1: str, pdb_path2: str) -> Optional[float]:
    pdb_parser = Bio.PDB.PDBParser(QUIET=True)
    model1 = next(iter(pdb_parser.get_structure("pdb1", pdb_path1)))
    model2 = next(iter(pdb_parser.get_structure("pdb2", pdb_path2)))
    chains1, chains2 = list(model1.get_chains()), list(model2.get_chains())
    assert len(chains1) == len(chains2) == 1, "can't extract if more than one chain"
    chain1, chain2 = chains1[0], chains2[0]

    chain1_res = [res for res in chain1.get_residues()]
    chain2_res = [res for res in chain2.get_residues()]

    chain1_ca = np.array([res["CA"].get_coord() for res in chain1_res])
    chain2_ca = np.array([res["CA"].get_coord() for res in chain2_res])

    close_residues = np.argwhere(scipy.spatial.distance.cdist(chain1_ca, chain2_ca) < INTERFACE_MIN_ATOM_DIST)
    if len(close_residues) == 0:
        print("Skipping transformation, missing interface between",
              os.path.basename(pdb_path1)[8:-4], os.path.basename(pdb_path2)[8:-4])
        return None

    chain1_interface, chain2_interface = set(), set()
    for i, j in close_residues:
        chain1_interface.add(i)
        chain2_interface.add(j)

    bfactors = [chain1_res[i]["CA"].get_bfactor() for i in chain1_interface] + \
               [chain2_res[i]["CA"].get_bfactor() for i in chain2_interface]
    return sum(bfactors) / len(bfactors)


def get_transformation_from_partials(partial_subunit1: PartialSubunit, partial_subunit2: PartialSubunit,

                                     representative_subunits_path: str, temp_folder: str,

                                     subunits_info: SubunitsInfo) -> Optional[TransformationInfo]:
    rep_struct1_path = extract_partial_from_representative(partial_subunit1, representative_subunits_path,
                                                           temp_folder, subunits_info)
    rep_struct2_path = extract_partial_from_representative(partial_subunit2, representative_subunits_path,
                                                           temp_folder, subunits_info)

    sample_struct1_path = os.path.join(temp_folder, f"sample1_{partial_subunit1.subunit_name}.pdb")
    extract_partial_subunit(partial_subunit1, sample_struct1_path)

    sample_struct2_path = os.path.join(temp_folder, f"sample2_{partial_subunit2.subunit_name}.pdb")
    extract_partial_subunit(partial_subunit2, sample_struct2_path)

    score = score_transformation(sample_struct1_path, sample_struct2_path)
    if score is None:
        return None

    af2trans_output = subprocess.check_output([AF2TRANS_BIN_PATH, rep_struct1_path, rep_struct2_path,
                                               sample_struct1_path, sample_struct2_path]).decode()
    assert af2trans_output.count(" | ") == 3, f"Unexpected output from AF2mer2trans {af2trans_output}"

    _, su1_desc, su2_desc, trans_nums = af2trans_output.split(" | ")
    rep_imposed_rmsds = (float(su1_desc.split("_")[0]), float(su2_desc.split("_")[0]))

    return TransformationInfo(
        subunit_names=(partial_subunit1.subunit_name, partial_subunit2.subunit_name),
        pdb_path=partial_subunit1.pdb_path,
        pdb_chain_ids=(partial_subunit1.chain_id, partial_subunit2.chain_id),
        transformation_numbers=trans_nums,
        rep_imposed_rmsds=rep_imposed_rmsds,
        score=score
    )


def get_pdb_to_partial_subunits(pdbs_folder: str, subunits_info: SubunitsInfo) -> Dict[str, List[PartialSubunit]]:
    pdb_path_to_partial_subunits: Dict[str, List[PartialSubunit]] = {}

    # for each pdb in folder
    for pdb_filename in os.listdir(pdbs_folder):
        if not pdb_filename.endswith(".pdb"):
            continue
        pdb_path = os.path.join(pdbs_folder, pdb_filename)

        partial_subunits: List[PartialSubunit] = []
        chain_to_seq = get_chain_to_seq(pdb_path, use_seqres=False)
        for chain_id, chain_seq in chain_to_seq.items():
            for subunit_info in subunits_info.values():
                subunit_seq = subunit_info.sequence
                if subunit_seq in chain_seq:
                    print(f"found full {subunit_info.name} in {pdb_filename} chain {chain_id}")
                    start_res_id = chain_seq.index(subunit_seq) + 1
                    end_res_id = start_res_id + len(subunit_seq) - 1
                    partial_subunits.append(PartialSubunit(subunit_name=subunit_info.name,
                                                           pdb_path=pdb_path,
                                                           chain_id=chain_id,
                                                           start_residue_id=start_res_id,
                                                           end_residue_id=end_res_id,
                                                           subunit_start_sequence_id=0,
                                                           is_complete=True)
                                            )
                elif chain_seq in subunit_seq:
                    start_residue_id = 1
                    end_residue_id = len(chain_seq)
                    print(f"found partial {subunit_info.name} in {pdb_filename} chain {chain_id}"
                          f"{(end_residue_id - start_residue_id + 1)}/{len(subunit_seq)}")
                    partial_subunits.append(PartialSubunit(subunit_name=subunit_info.name,
                                                           pdb_path=pdb_path,
                                                           chain_id=chain_id,
                                                           start_residue_id=start_residue_id,
                                                           end_residue_id=end_residue_id,
                                                           subunit_start_sequence_id=subunit_seq.index(chain_seq),
                                                           is_complete=False)
                                            )
                else:
                    min_match_length = 10
                    start_ind = 0
                    while start_ind < len(chain_seq) - min_match_length:
                        if chain_seq[start_ind] == "X":
                            start_ind += 1
                            continue
                        end_ind = start_ind + min_match_length
                        while len(chain_seq[start_ind:end_ind].replace("X", "")) < min_match_length:
                            end_ind += 1
                            if end_ind >= len(chain_seq):
                                end_ind += 1
                                break
                        if end_ind >= len(chain_seq):
                            break

                        if chain_seq[start_ind:end_ind] not in subunit_seq:
                            start_ind += 1
                            continue
                        while end_ind < len(chain_seq) and chain_seq[start_ind:end_ind] in subunit_seq:
                            end_ind += 1
                        if chain_seq[start_ind:end_ind] not in subunit_seq:
                            end_ind -= 1

                        start_residue_id = start_ind + 1  # get_chain_seq function is 0-based, while res_id is 1-based
                        end_residue_id = end_ind + 1

                        subunit_start_ind = subunit_seq.index(chain_seq[start_ind:end_ind])
                        print(f"found small partial {subunit_info.name} in {pdb_filename} chain {chain_id} "
                              f"starting at index {start_ind} to {end_ind} (on subunit {subunit_start_ind})"
                              f"length {(end_ind - start_ind)}/{len(subunit_seq)}")
                        partial_subunits.append(PartialSubunit(subunit_name=subunit_info.name,
                                                               pdb_path=pdb_path,
                                                               chain_id=chain_id,
                                                               start_residue_id=start_residue_id,
                                                               end_residue_id=end_residue_id - 1,  # inclusive index
                                                               subunit_start_sequence_id=subunit_start_ind,
                                                               is_complete=False)
                                                )
                        start_ind = end_ind

        # print(f"found {len(partial_subunits)} partial subunits in {pdb_filename}")
        partial_subunits = sorted(partial_subunits, key=lambda x: (x.subunit_name, x.chain_id, x.start_residue_id))
        pdb_path_to_partial_subunits[pdb_path] = partial_subunits
    return pdb_path_to_partial_subunits


def extract_representative_subunits(pdb_path_to_partial_subunits: Dict[str, List[PartialSubunit]],

                                    subunits_info: SubunitsInfo, representative_subunits_path: str):
    rep_structs: Dict[SubunitName, Tuple[float, PartialSubunit]] = {}
    for pdb_path, partial_subunits in pdb_path_to_partial_subunits.items():
        for partial_subunit in partial_subunits:
            if not partial_subunit.is_complete:
                continue
            subunit_residues = _get_partial_subunit_residues(partial_subunit)
            plddt_score = sum([res["CA"].get_bfactor() for res in subunit_residues]) / len(subunit_residues)
            if rep_structs.get(partial_subunit.subunit_name, (-1, None))[0] < plddt_score:
                rep_structs[partial_subunit.subunit_name] = (plddt_score, partial_subunit)
    assert len(rep_structs) == len(subunits_info), "missing rep subunits for" + \
                                                   str(set(subunits_info.keys()) - set(rep_structs.keys()))
    for subunit_name, (plddt_score, partial_subunit) in rep_structs.items():
        print(f"rep {subunit_name} has plddt score {plddt_score}")
        subunit_info = subunits_info[subunit_name]
        rep_struct_path = os.path.join(representative_subunits_path, f"{subunit_name}.pdb")
        print("extracting partial subunit to", partial_subunit)
        extract_partial_subunit(partial_subunit, rep_struct_path)
        copy_pdb_set_start_offset(rep_struct_path, subunit_info.start_res, rep_struct_path)
        for chain_name, ident_subunit_name in zip(subunit_info.chain_names, subunit_info.get_chained_names()):
            copy_pdb_rename_chain(rep_struct_path, chain_name,
                                  os.path.join(representative_subunits_path, f"{ident_subunit_name}.pdb"))
        os.remove(rep_struct_path)


def extract_transformations(pdb_path_to_partial_subunits: Dict[str, List[PartialSubunit]], subunits_info: SubunitsInfo,

                            representative_subunits_path: str, transformations_path: str):
    temp_folder = os.path.join(transformations_path, "temp_transformations")
    transformations_by_pdb_path: Dict[str, List[TransformationInfo]] = {}
    for pdb_path, partial_subunits in pdb_path_to_partial_subunits.items():
        print("- Extracting pairwise transformations from file", pdb_path)
        if os.path.exists(temp_folder):
            print("removing temp folder")
            shutil.rmtree(temp_folder)
        os.makedirs(temp_folder)
        transformations_by_pdb_path[pdb_path] = []
        for partial_subunit_ind_i in range(len(partial_subunits)):
            partial_subunit1 = partial_subunits[partial_subunit_ind_i]

            for partial_subunit_ind_j in range(partial_subunit_ind_i + 1, len(partial_subunits)):
                partial_subunit2 = partial_subunits[partial_subunit_ind_j]

                transformation_info = get_transformation_from_partials(partial_subunit1, partial_subunit2,
                                                                       representative_subunits_path, temp_folder,
                                                                       subunits_info)
                if transformation_info is not None:
                    transformations_by_pdb_path[pdb_path].append(transformation_info)
        shutil.rmtree(temp_folder)

    transformations_by_subunit_pair: Dict[Tuple[str, str], List[TransformationInfo]] = defaultdict(list)
    for pdb_path, transformations in transformations_by_pdb_path.items():
        for transformation in transformations:
            transformations_by_subunit_pair[transformation.subunit_names].append(transformation)

    for (subunit_name1, subunit_name2), transformations in transformations_by_subunit_pair.items():
        print(f"found {len(transformations)} transformations between {subunit_name1} and {subunit_name2}")
        transformations = sorted(transformations, key=lambda x: x.score, reverse=True)
        output_file_path = os.path.join(transformations_path, f"tmp_{subunit_name1}_plus_{subunit_name2}")

        with open(output_file_path, "w") as f:
            for i, transformation in enumerate(transformations):
                description = f"{transformation.rep_imposed_rmsds[0]}_{transformation.rep_imposed_rmsds[1]}_" \
                              f"{transformation.pdb_chain_ids[0]}_{transformation.pdb_chain_ids[1]}_" \
                              f"{os.path.basename(transformation.pdb_path)}"
                f.write(f"{i + 1} | {transformation.score} | {description} | {transformation.transformation_numbers}\n")

        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_path, f"{aliases_c1[c1]}_plus_{aliases_c2[c2]}")
                shutil.copy(output_file_path, alias_output_path)
        os.remove(output_file_path)


def run_combfold(representative_subunits_path: str, subunits_info: SubunitsInfo, transformations_path: str,

                 crosslinks_path: Optional[str], output_path: str, output_cif: bool = False,

                 max_results_number: int = 5, subunits_group1: Optional[List[str]] = None,):
    # prepare and run assembly
    with open(os.path.join(representative_subunits_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:
            if subunits_group1 is not None and chained_subunit_name in subunits_group1:
                f.write(f"{chained_subunit_name}.pdb 1\n")
            else:
                f.write(f"{chained_subunit_name}.pdb\n")
    os.chdir(representative_subunits_path)
    if crosslinks_path is not None:
        shutil.copy(crosslinks_path, os.path.join(representative_subunits_path, "xlink_consts.txt"))
    else:
        open("xlink_consts.txt", "w").close()

    subprocess.run(f"{COMB_ASSEMBLY_BIN_PATH} chain.list {transformations_path}/ 900 100 xlink_consts.txt "
                   f"-b 0.05 -t 80 > output.log 2>&1", shell=True)
    print("--- Finished combinatorial assembly, writing output models")

    # build pdbs from assembly output
    clusters_path = os.path.join(representative_subunits_path, "output_clustered.res")
    if not os.path.exists(clusters_path):
        print(f"Could not assemble, exiting")
        return
    assembled_files = create_complexes(clusters_path, first_result=0, last_result=max_results_number,
                                       output_folder=os.path.join(output_path, "assembled_results"),
                                       output_cif=output_cif)

    confidence = []
    for result_as_str in open(clusters_path, "r").read().split("\n")[:len(assembled_files)]:
        if not result_as_str.strip():
            continue
        splitted_result = result_as_str.split(" ")
        confidence.append(float(splitted_result[splitted_result.index("weightedTransScore") + 1]))

    with open(os.path.join(output_path, "assembled_results", "confidence.txt"), "w") as f:
        for filename, c in zip(assembled_files, confidence):
            f.write(f"{filename} {c}\n")

    print(f"--- Assembled {len(assembled_files)} complexes, confidence: {min(confidence)}-{max(confidence)}")


def run_on_pdbs_folder(subunits_json_path: str, pdbs_folder: str, output_path: str,

                       crosslinks_path: Optional[str] = None, output_cif: bool = False, max_results_number: int = 5):
    pdbs_folder = os.path.abspath(pdbs_folder)
    output_path = os.path.abspath(output_path)

    if not os.path.exists(COMB_ASSEMBLY_BIN_PATH):
        print(f"combinatorial assembly binary not found at {COMB_ASSEMBLY_BIN_PATH}, compile it by: \n"
              f"cd {os.path.dirname(COMB_ASSEMBLY_BIN_PATH)} && make")
        return

    if os.path.exists(output_path) and os.listdir(output_path):
        print(f"output path {output_path} is not empty, exiting")
        return

    subunits_info: SubunitsInfo = read_subunits_info(subunits_json_path)

    # representative_subunits_path is also the assembly algorithm output path
    representative_subunits_path = os.path.join(output_path, "_unified_representation", "assembly_output")
    transformations_path = os.path.join(output_path, "_unified_representation", "transformations")
    os.makedirs(representative_subunits_path, exist_ok=True)
    os.makedirs(transformations_path, exist_ok=True)

    print("--- Searching for subunits in supplied PDB files")
    pdb_path_to_partial_subunits = get_pdb_to_partial_subunits(pdbs_folder, subunits_info)

    print("--- Extracting representative subunits (for each subunit, its best scored model in the PDBs folder)")
    extract_representative_subunits(pdb_path_to_partial_subunits, subunits_info, representative_subunits_path)

    print("--- Extracting pairwise transformations between subunits (from each PDB file with 2 or more subunits)")
    extract_transformations(pdb_path_to_partial_subunits, subunits_info, representative_subunits_path,
                            transformations_path)

    print("--- Finished building unified representation")

    print("--- Running combinatorial assembly algorithm, may take a while")
    run_combfold(representative_subunits_path, subunits_info, transformations_path, crosslinks_path, output_path,
                 output_cif, max_results_number)


if __name__ == '__main__':
    if len(sys.argv) == 4:
        run_on_pdbs_folder(os.path.abspath(sys.argv[1]), os.path.abspath(sys.argv[2]), os.path.abspath(sys.argv[3]))
    elif len(sys.argv) == 5:
        run_on_pdbs_folder(os.path.abspath(sys.argv[1]), os.path.abspath(sys.argv[2]), os.path.abspath(sys.argv[3]),
                           crosslinks_path=os.path.abspath(sys.argv[4]))
    else:
        print("usage: <script> subunits_info pdbs_folder output_path <optional: crosslinks.txt>")