From f087c59458f2f287c8757540b3073ed5b08e54fb Mon Sep 17 00:00:00 2001 From: Gyuri Kim Date: Sat, 8 Aug 2026 13:43:22 +0900 Subject: [PATCH] add ipsae and pdockq2 in final prediction stage --- colabfold/alphafold/ipsae.py | 98 ++++++++++++++++++++++++++++++++++++ colabfold/batch.py | 20 +++++++- tests/test_colabfold.py | 8 +-- 3 files changed, 121 insertions(+), 5 deletions(-) create mode 100644 colabfold/alphafold/ipsae.py diff --git a/colabfold/alphafold/ipsae.py b/colabfold/alphafold/ipsae.py new file mode 100644 index 000000000..1b4ed8d45 --- /dev/null +++ b/colabfold/alphafold/ipsae.py @@ -0,0 +1,98 @@ +"""Interface confidence scores.""" + +# Reimplemented from https://github.com/DunbrackLab/IPSAE (MIT). + +import numpy as np + +from alphafold.common import residue_constants +from alphafold.common.protein import PDB_CHAIN_IDS + +# Angstrom. +DEFAULT_PAE_CUTOFF = 15.0 +PDOCKQ_DIST_CUTOFF = 8.0 + + +def _ptm_score(pae, d0): + """TM-score kernel.""" + return 1.0 / (1.0 + (pae / d0) ** 2.0) + + +def _calc_d0_array(n_res): + """TM-score d0.""" + length = np.maximum(26.0, np.asarray(n_res, dtype=np.float64)) + return np.maximum(1.0, 1.24 * (length - 15.0) ** (1.0 / 3.0) - 1.8) + + +def get_interface_scores(pae, plddt, asym_id, atom_positions, atom_mask, + pae_cutoff=DEFAULT_PAE_CUTOFF): + """Return ipSAE, pDockQ and pDockQ2 per chain pair.""" + asym_id = np.asarray(asym_id) + unique_chains = np.unique(asym_id) + if len(unique_chains) < 2: + return {} + + pae = np.asarray(pae, dtype=np.float64) + plddt = np.asarray(plddt, dtype=np.float64) + atom_positions = np.asarray(atom_positions, dtype=np.float64) + + # Use CA for glycine. + atom_mask = np.asarray(atom_mask) + ca_idx = residue_constants.atom_order["CA"] + cb_idx = residue_constants.atom_order["CB"] + has_cb = atom_mask[:, cb_idx] > 0.5 + cb_coords = np.where(has_cb[:, None], + atom_positions[:, cb_idx], atom_positions[:, ca_idx]) + distances = np.sqrt(((cb_coords[:, None] - cb_coords[None, :]) ** 2).sum(-1)) + # Exclude unplaced residues. + unplaced = atom_mask[:, ca_idx] <= 0.5 + distances[unplaced, :] = np.inf + distances[:, unplaced] = np.inf + + chain_label = {chain: PDB_CHAIN_IDS[int(chain)] for chain in unique_chains} + ipsae, pdockq, pdockq2 = {}, {}, {} + for chain_1 in unique_chains: + for chain_2 in unique_chains: + if chain_1 == chain_2: + continue + key = f"{chain_label[chain_1]}-{chain_label[chain_2]}" + pair_mask = np.outer(asym_id == chain_1, asym_id == chain_2) + + # ipSAE. + valid_pairs = pair_mask & (pae < pae_cutoff) + n0_res = valid_pairs.sum(axis=1) + d0_res = _calc_d0_array(n0_res) + ptm_sums = (_ptm_score(pae, d0_res[:, None]) * valid_pairs).sum(axis=1) + ipsae_by_res = np.divide(ptm_sums, n0_res, + out=np.zeros_like(ptm_sums), where=n0_res > 0) + ipsae[key] = round(float(ipsae_by_res.max()), 6) + + # pDockQ and pDockQ2. + contacts = pair_mask & (distances <= PDOCKQ_DIST_CUTOFF) + n_contacts = contacts.sum() + if n_contacts > 0: + interface_res = contacts.any(axis=1) | contacts.any(axis=0) + mean_plddt = plddt[interface_res].mean() + pdockq_val = 0.724 / (1.0 + np.exp( + -0.052 * (mean_plddt * np.log10(n_contacts) - 152.611))) + 0.018 + mean_ptm = _ptm_score(pae[contacts], 10.0).mean() + pdockq2_val = 1.31 / (1.0 + np.exp( + -0.075 * (mean_plddt * mean_ptm - 84.733))) + 0.005 + else: + pdockq_val, pdockq2_val = 0.0, 0.0 + if chain_1 < chain_2: # Symmetric. + pdockq[key] = round(float(pdockq_val), 4) + pdockq2[key] = round(float(pdockq2_val), 4) + + return {"ipsae": ipsae, "pdockq": pdockq, "pdockq2": pdockq2} + + +def format_ipsae(ipsae_scores): + """Format per-interface maxima.""" + pair_max = {} + for key, value in ipsae_scores.items(): + chain_1, chain_2 = key.split("-") + pair = key if chain_1 < chain_2 else f"{chain_2}-{chain_1}" + pair_max[pair] = max(pair_max.get(pair, 0.0), value) + if len(pair_max) == 1: + return f"{next(iter(pair_max.values())):.3g}" + return ",".join(f"{pair}:{value:.3g}" for pair, value in sorted(pair_max.items())) diff --git a/colabfold/batch.py b/colabfold/batch.py index de08295c6..36b837f44 100644 --- a/colabfold/batch.py +++ b/colabfold/batch.py @@ -79,7 +79,7 @@ pdb_to_string, ) from colabfold.relax import relax_me -from colabfold.alphafold import extra_ptm +from colabfold.alphafold import extra_ptm, ipsae from Bio.PDB import MMCIFParser, PDBParser, MMCIF2Dict from Bio.PDB.PDBIO import Select @@ -587,6 +587,24 @@ def callback(result, recycles): for k in ["ptm", "iptm"]: if k in conf[-1]: scores[k] = np.around(conf[-1][k], 2).item() + if is_complex: + try: + asym_id = input_features["asym_id"] + if asym_id.ndim > 1: asym_id = asym_id[0] + interface_scores = ipsae.get_interface_scores( + pae=pae, + plddt=plddt, + asym_id=asym_id[:seq_len], + atom_positions=result["structure_module"]["final_atom_positions"][:seq_len], + atom_mask=result["structure_module"]["final_atom_mask"][:seq_len]) + scores.update(interface_scores) + if interface_scores: + conf[-1]["print_line"] += ( + f" ipSAE={ipsae.format_ipsae(interface_scores['ipsae'])}" + f" pDockQ2={ipsae.format_ipsae(interface_scores['pdockq2'])}" + ) + except Exception as e: + logger.warning(f"Could not compute ipSAE/pDockQ interface scores: {e}") del pae del plddt file = files.get("scores", "json") diff --git a/tests/test_colabfold.py b/tests/test_colabfold.py index 1a1ed423c..667dfaf8b 100644 --- a/tests/test_colabfold.py +++ b/tests/test_colabfold.py @@ -221,7 +221,7 @@ def test_complex(pytestconfig, caplog, tmp_path, prediction_test): 'Setting max_seq=252, max_extra_seq=1152', 'alphafold2_multimer_v1_model_1_seed_000 took 0.0s (3 recycles)', "reranking models by 'multimer' metric", - 'rank_001_alphafold2_multimer_v1_model_1_seed_000 pLDDT=94.4 pTM=0.884 ipTM=0.878', + 'rank_001_alphafold2_multimer_v1_model_1_seed_000 pLDDT=94.4 pTM=0.884 ipTM=0.878 ipSAE=0.787 pDockQ2=0.905', 'Done' ] for x in expected: @@ -260,7 +260,7 @@ def test_complex_ptm(pytestconfig, caplog, tmp_path, prediction_test): 'Setting max_seq=512, max_extra_seq=5120', 'alphafold2_ptm_model_1_seed_000 took 0.0s (3 recycles)', "reranking models by 'multimer' metric", - 'rank_001_alphafold2_ptm_model_1_seed_000 pLDDT=92 pTM=0.846 ipTM=0.849', + 'rank_001_alphafold2_ptm_model_1_seed_000 pLDDT=92 pTM=0.846 ipTM=0.849 ipSAE=0.741 pDockQ2=0.863', 'Done' ] for x in expected: @@ -300,7 +300,7 @@ def test_complex_monomer_ptm(pytestconfig, caplog, tmp_path, prediction_test): 'Setting max_seq=512, max_extra_seq=5120', 'alphafold2_ptm_model_1_seed_000 took 0.0s (3 recycles)', "reranking models by 'multimer' metric", - 'rank_001_alphafold2_ptm_model_1_seed_000 pLDDT=95.6 pTM=0.867 ipTM=0.864', + 'rank_001_alphafold2_ptm_model_1_seed_000 pLDDT=95.6 pTM=0.867 ipTM=0.864 ipSAE=0.739 pDockQ2=0.934', 'Done' ] for x in expected: @@ -340,7 +340,7 @@ def test_complex_monomer(pytestconfig, caplog, tmp_path, prediction_test): 'Setting max_seq=252, max_extra_seq=1152', 'alphafold2_multimer_v1_model_1_seed_000 took 0.0s (3 recycles)', "reranking models by 'multimer' metric", - 'rank_001_alphafold2_multimer_v1_model_1_seed_000 pLDDT=95.3 pTM=0.866 ipTM=0.861', + 'rank_001_alphafold2_multimer_v1_model_1_seed_000 pLDDT=95.3 pTM=0.866 ipTM=0.861 ipSAE=0.714 pDockQ2=0.932', 'Done' ] for x in expected: