Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 98 additions & 0 deletions colabfold/alphafold/ipsae.py
Original file line number Diff line number Diff line change
@@ -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()))
20 changes: 19 additions & 1 deletion colabfold/batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
8 changes: 4 additions & 4 deletions tests/test_colabfold.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
Loading