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
7 changes: 6 additions & 1 deletion bionemo_ir/data/schemas/basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -833,10 +833,15 @@ def get_scores(self) -> dict:
iptm = float(self["iptm"])
if self["max_pae"] is not None and not np.isnan(self["max_pae"]):
max_pae = float(self["max_pae"])
return {
scores = {
"plddt": plddt,
"ptm": ptm,
"iptm": iptm,
"pae": pae,
"max_pae": max_pae,
}
# Optional scalar metadata set by a model post-processor, e.g. the
# aligned-token mask convention behind ptm / iptm (OpenFold3).
if self.get("ptm_frame_mask") is not None:
scores["ptm_frame_mask"] = str(self["ptm_frame_mask"])

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This return is redundant, we should leave to the original.

return scores
117 changes: 100 additions & 17 deletions bionemo_ir/pipeline/models/openfold3/postprocessor.py

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Re-check with OSS and AF3 and confirm this is the real bug for computing scores.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This below is recommending fixes with efficient memory, basically we should reuse atomize_utils.py::get_token_frame_mask under the model-level:

diff --git a/bionemo_ir/_torch/modules/openfold3/utils/atomize_utils.py b/bionemo_ir/_torch/modules/openfold3/utils/atomize_utils.py
index 2a11015..529a919 100644
--- a/bionemo_ir/_torch/modules/openfold3/utils/atomize_utils.py
+++ b/bionemo_ir/_torch/modules/openfold3/utils/atomize_utils.py
@@ -295,21 +295,87 @@ def get_token_atom_index_offset(atom_name: str, restype: torch.Tensor):
     return token_atom_index_offset, token_atom_mask
 
 
-def get_token_frame_atoms(
+def _closest_atoms_to_start_atoms(
+    x: torch.Tensor,
+    atom_mask: torch.Tensor,
+    atom_asym_id: torch.Tensor,
+    start_atom_index: torch.Tensor,
+    eps: float,
+    inf: float,
+) -> tuple[torch.Tensor, torch.Tensor]:
+    """Indices of the two closest same-chain atoms to each token's start atom.
+
+    A dense neighbour search over all atom pairs computes ``N_atom`` rows of which
+    only the ``N_token`` start-atom rows are ever read, so the query axis here is the
+    tokens. Distances and the pair mask are formed exactly as the dense search forms
+    them, so the selected indices are unchanged, but the working set is
+    ``N_token x N_atom`` per diffusion sample rather than ``N_atom^2`` -- the
+    difference between single-digit GB and tens of GB on a large complex.
+
+    Args:
+        x:
+            [*, N_atom, 3] Atom positions
+        atom_mask:
+            [*, N_atom] Atom mask
+        atom_asym_id:
+            [*, N_atom] Chain index, broadcast to atoms
+        start_atom_index:
+            [*, N_token] Index of the first atom of each token
+        eps:
+            Small constant for numerical stability
+        inf:
+            Large constant for numerical stability
+
+    Returns:
+        ([*, N_token], [*, N_token])
+        Indices of the closest and second closest atom to each token's start atom
+    """
+    leading_shape = x.shape[:-2]
+    atom_mask = torch.broadcast_to(atom_mask, (*leading_shape, atom_mask.shape[-1]))
+    atom_asym_id = torch.broadcast_to(atom_asym_id, (*leading_shape, atom_asym_id.shape[-1]))
+
+    # Position, mask and chain of the query (start) atoms
+    start_x = torch.gather(x, dim=-2, index=start_atom_index.unsqueeze(-1).expand(*start_atom_index.shape, 3))
+    start_atom_mask = torch.gather(atom_mask, dim=-1, index=start_atom_index)
+    start_asym_id = torch.gather(atom_asym_id, dim=-1, index=start_atom_index)
+
+    # Pairwise mask over (start atom, atom): both present, and within the same chain
+    # [*, N_token, N_atom]
+    pair_mask = start_atom_mask[..., None] * atom_mask[..., None, :]
+    pair_mask = pair_mask * (start_asym_id[..., None] == atom_asym_id[..., None, :])
+
+    # Distance from every start atom to every atom
+    # [*, N_token, N_atom]
+    d = torch.sum(eps + (start_x[..., None, :] - x[..., None, :, :]) ** 2, dim=-1) ** 0.5
+    d = d * pair_mask + inf * (1 - pair_mask)
+
+    # Index 0 is the start atom itself, so 1 and 2 are its two closest neighbours
+    _, closest_atom_index = torch.topk(d, k=3, dim=-1, largest=False)
+    return closest_atom_index[..., 1], closest_atom_index[..., 2]
+
+
+def get_token_frame_mask(
     batch: dict,
     x: torch.Tensor,
     atom_mask: torch.Tensor,
     angle_threshold: float = 25.0,
     eps: float = 1e-8,
     inf: float = 1e9,
-):
+) -> torch.Tensor:
     """
-    Extract frame atoms per token, which returns
+    Mask of tokens whose frame is valid, from the frame atoms
         -   (N, Ca, C) for standard amino acid residues
         -   (C3', C1', C4') for standard nucleotide residues
         -   closest neighbors for atomized tokens (modified residues and ligands),
             subject to additional angle and chain constraints from Subsection 4.3.2
 
+    A frame is valid when its three atoms are present, lie in one chain, and -- for
+    atomized tokens, whose frame comes from nearest neighbours rather than a known
+    backbone -- span an angle within ``angle_threshold`` of neither 0 nor 180 degrees.
+    This is the ``has_frame`` input of the pTM / ipTM outer maximum (AF3 SI 5.9.1):
+    only a token with a frame can be the aligned token. The frame atom positions
+    themselves are used only to test that angle and are not returned.
+
     Args:
         batch:
             Feature dictionary
@@ -324,37 +390,28 @@ def get_token_frame_atoms(
         inf:
             Large constant for numerical stability
     Returns:
-        phi:
-            ([*, N_token, 3], [*, N_token, 3], [*, N_token, 3])
-            Tuple of three frame atoms
         valid_frame_mask:
             [*, N_token] Mask denoting valid frames
     """
-    # Create pairwise atom mask
-    pair_mask = atom_mask[..., None] * atom_mask[..., None, :]
-
-    # Update pairwise atom mask
-    # Restrict to atoms within the same chain
+    # Chain index per atom, to restrict frames to atoms within the same chain
     atom_asym_id = broadcast_token_feat_to_atoms(
         token_mask=batch["token_mask"],
         num_atoms_per_token=batch["num_atoms_per_token"],
         token_feat=batch["asym_id"],
     )
-    atom_asym_id_mask = atom_asym_id[..., None] == atom_asym_id[..., None, :]
-    pair_mask = pair_mask * atom_asym_id_mask
-
-    # Compute distance matrix
-    # [*, N_atom, N_atom]
-    d = torch.sum(eps + (x[..., None, :] - x[..., None, :, :]) ** 2, dim=-1) ** 0.5
-    d = d * pair_mask + inf * (1 - pair_mask)
 
     # Find indices of two closest atoms for start atoms
     # [*, N_token]
     start_atom_index = batch["start_atom_index"].long()
     start_atom_index = start_atom_index.expand(*x.shape[:-2], start_atom_index.shape[-1])
-    _, closest_atom_index = torch.topk(d, k=3, dim=-1, largest=False)
-    a_index = torch.gather(closest_atom_index[..., 1], dim=-1, index=start_atom_index)
-    c_index = torch.gather(closest_atom_index[..., 2], dim=-1, index=start_atom_index)
+    a_index, c_index = _closest_atoms_to_start_atoms(
+        x=x,
+        atom_mask=atom_mask,
+        atom_asym_id=atom_asym_id,
+        start_atom_index=start_atom_index,
+        eps=eps,
+        inf=inf,
+    )
 
     # Construct indices of atoms used for frame construction
     # [*, N_token]
@@ -401,7 +458,8 @@ def get_token_frame_atoms(
         },
     }
 
-    # Extract coordinates
+    # Extract chain, presence and coordinates of each frame atom. The coordinates
+    # serve only the angle test below; they are not part of the result.
     for key in frame_atoms:
         frame_atoms[key].update(
             {
@@ -457,11 +515,4 @@ def get_token_frame_atoms(
     )
 
     # Compute final valid frame mask
-    valid_frame_mask = valid_frame_mask_angle * valid_frame_mask_atom * valid_frame_mask_asym_id
-    phi = (
-        frame_atoms["a"]["atom_positions"],
-        frame_atoms["b"]["atom_positions"],
-        frame_atoms["c"]["atom_positions"],
-    )
-
-    return phi, valid_frame_mask
+    return valid_frame_mask_angle * valid_frame_mask_atom * valid_frame_mask_asym_id

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

At the confidence.py, the computation of valid_frame_mask from get_token_frame_atoms:

diff --git a/bionemo_ir/_torch/modules/openfold3/confidence.py b/bionemo_ir/_torch/modules/openfold3/confidence.py
index 6b8c081..98e9f58 100644
--- a/bionemo_ir/_torch/modules/openfold3/confidence.py
+++ b/bionemo_ir/_torch/modules/openfold3/confidence.py
@@ -22,6 +22,7 @@ from bionemo_ir._torch.layers.linear import Linear
 from bionemo_ir._torch.layers.transformers.pairformer import PairformerModule
 from bionemo_ir._torch.modules.openfold3.utils.atomize_utils import (
     broadcast_token_feat_to_atoms,
+    get_token_frame_mask,
     get_token_representative_atoms,
     max_atom_per_token_masked_select,
 )
@@ -657,6 +658,9 @@ class AuxiliaryHeadsAllAtom(nn.Module):
                         Predicted binned PLDDT logits
                     "pae_logits" ([*, N_token, N_token, 64]):
                         Predicted binned PAE logits
+                    "valid_frame_mask" ([*, N_token]):
+                        Tokens with a valid frame, the ``has_frame`` input of the
+                        pTM / ipTM outer maximum. Present with "pae_logits" only.
                     "pde_logits" ([*, N_token, N_token, 64]):
                         Predicted binned PDE logits
                     "experimentally_resolved_logits" ([*, N_atom, 2]):
@@ -721,6 +725,12 @@ class AuxiliaryHeadsAllAtom(nn.Module):
 
         if self.config.pae.enabled:
             aux_out["pae_logits"] = self.pae(zij).to(device=out_device)
+            # has_frame for the pTM / ipTM outer maximum, from the sampled
+            # coordinates, so it is per diffusion sample. Only the PAE head feeds
+            # pTM / ipTM, so nothing needs it when that head is off.
+            aux_out["valid_frame_mask"] = get_token_frame_mask(
+                batch=batch, x=atom_positions_predicted, atom_mask=batch["atom_mask"]
+            ).to(device=out_device)
 
         aux_out["pde_logits"] = pde_logits.to(device=out_device)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

And below suggestions on the file postprocess.py will move tensors computation for scores on the GPU devices.

Original file line number Diff line number Diff line change
Expand Up @@ -121,8 +121,9 @@ def __call__(

# --- Confidence scores from logits ---
plddt = _compute_plddt(output, best_idx, n_tokens, atom_to_token, atom_mask_bool)
ptm = _compute_ptm(output, best_idx, n_tokens)
iptm = _compute_iptm(output, best_idx, n_tokens, chain_indices)
has_frame = _aligned_token_mask(batch, n_tokens)
ptm = _compute_ptm(output, best_idx, n_tokens, has_frame=has_frame)
iptm = _compute_iptm(output, best_idx, n_tokens, chain_indices, has_frame=has_frame)
pae = _compute_pae(output, best_idx, n_tokens)
Comment on lines -124 to 127

@ducta3141 ducta3141 Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

        pae_logits = _pae_logits(output, best_idx, n_tokens)
        has_frame = _frame_mask(output, best_idx, n_tokens)
        ptm = _compute_ptm(pae_logits, n_tokens, has_frame=has_frame)
        iptm = _compute_iptm(pae_logits, n_tokens, chain_indices, has_frame=has_frame)
        pae = _compute_pae(pae_logits)

max_pae = float(np.max(pae)) if pae is not None else None

Expand Down Expand Up @@ -160,6 +161,10 @@ def __call__(
residue_names=residue_names,
mol_types=mol_types_out,
)
# Convention of the aligned-token (frame) mask behind ptm / iptm, carried
# into get_scores(): "polymer_tokens" = interim ~is_atomized mask,
# "none" = every token eligible (is_atomized absent from the batch).
result["ptm_frame_mask"] = "none" if has_frame is None else "polymer_tokens"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove this part.


return result

Expand All @@ -175,6 +180,37 @@ def _cpu(t: Any) -> torch.Tensor:
return torch.as_tensor(t)


def _bin_centers(n_bins: int, bin_min: float, bin_max: float) -> torch.Tensor:
"""Midpoints of ``n_bins`` equal-width bins on ``[bin_min, bin_max]``.

Matches AF3 and upstream OpenFold3 (``openfold3/core/metrics/confidence.py::
get_bin_centers``): 0.25, 0.75, ..., 31.75 Å for the 64-bin PAE head and
0.01, 0.03, ..., 0.99 for the 50-bin pLDDT head -- not the end-point-inclusive
positions of ``torch.linspace(bin_min, bin_max, n_bins)``.
"""
width = (bin_max - bin_min) / n_bins
return bin_min + width * (torch.arange(n_bins, dtype=torch.float32) + 0.5)


def _aligned_token_mask(batch: dict[str, Any], n_tokens: int) -> torch.Tensor | None:
"""Tokens eligible as the aligned token ``i`` in the pTM / ipTM max (``has_frame``).

Interim stand-in for the coordinate-based frame validity of upstream OpenFold3
(``openfold3/core/utils/atomize_utils.py::get_token_frame_atoms``): polymer
tokens are eligible; atomized tokens (``batch["is_atomized"]``: ligand atoms,
ions) are scored as ``j`` but never used as aligned tokens. Upstream additionally
admits atomized tokens whose nearest-neighbour local frame is valid and requires
the backbone frame atoms of polymer residues to be present. Returns ``None``
(every token eligible) only when the flag is absent from the batch; an input
without polymer tokens (ligand-only query) yields an all-False mask, for which
pTM / ipTM are reported as NaN (see ``_tm_score_from_pae_logits``).
"""
is_atomized = batch.get("is_atomized")
if is_atomized is None:
return None
return ~_cpu(is_atomized).reshape(-1)[:n_tokens].bool()
Comment on lines +195 to +211

@ducta3141 ducta3141 Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
def _aligned_token_mask(batch: dict[str, Any], n_tokens: int) -> torch.Tensor | None:
"""Tokens eligible as the aligned token ``i`` in the pTM / ipTM max (``has_frame``).
Interim stand-in for the coordinate-based frame validity of upstream OpenFold3
(``openfold3/core/utils/atomize_utils.py::get_token_frame_atoms``): polymer
tokens are eligible; atomized tokens (``batch["is_atomized"]``: ligand atoms,
ions) are scored as ``j`` but never used as aligned tokens. Upstream additionally
admits atomized tokens whose nearest-neighbour local frame is valid and requires
the backbone frame atoms of polymer residues to be present. Returns ``None``
(every token eligible) only when the flag is absent from the batch; an input
without polymer tokens (ligand-only query) yields an all-False mask, for which
pTM / ipTM are reported as NaN (see ``_tm_score_from_pae_logits``).
"""
is_atomized = batch.get("is_atomized")
if is_atomized is None:
return None
return ~_cpu(is_atomized).reshape(-1)[:n_tokens].bool()
def _frame_mask(output: dict, best_idx: int, n_tokens: int) -> torch.Tensor | None:
"""``has_frame`` for the pTM / ipTM maximum, as emitted by the confidence head.
The head computes it from the sampled coordinates, so it carries the diffusion
sample axis and arrives as 0/1 floats in the output dtype. Absent when the PAE
head is disabled, in which case there is no pTM to restrict. Stays on its
original device, alongside the logits it will gate.
"""
mask = output.get("valid_frame_mask")
if mask is None:
return None
mask = torch.as_tensor(mask)
if mask.dim() == 3:
mask = mask[0, best_idx]
elif mask.dim() == 2:
mask = mask[0]
return mask[:n_tokens].bool()
def _pae_logits(output: dict, best_idx: int, n_tokens: int) -> torch.Tensor | None:
"""The selected sample's PAE logits, cropped to the real tokens.
Left on its original device: ``torch.as_tensor`` preserves it, unlike
``_cpu``. The engine calls the post-processor with the model's own output, so
this is normally GPU memory, and every consumer reduces it.
"""
logits = output.get("pae_logits")
if logits is None:
return None
logits = torch.as_tensor(logits)
if logits.dim() == 5:
logits = logits[0, best_idx] # (N_tokens, N_tokens, n_bins)
elif logits.dim() == 4:
logits = logits[0]
return logits[:n_tokens, :n_tokens]
def _plddt_per_atom(output: dict) -> torch.Tensor | None:
"""Per-atom pLDDT on a 0-100 scale, as ``(n_samples, N_atom)``.
Both the diffusion-sample choice and the reported per-token score are this
same expectation, so it runs once here instead of once in each. Stays on the
logits' device; only the selected row is ever copied to the host.
"""
logits = output.get("plddt_logits")
if logits is None:
return None
logits = torch.as_tensor(logits)
if logits.dim() == 4:
logits = logits[0] # (B, S, N_atom, n_bins) -> (S, N_atom, n_bins)
elif logits.dim() == 3:
logits = logits[:1] # (B, N_atom, n_bins) -> a single sample
if logits.dim() == 2:
logits = logits.unsqueeze(0) # (N_atom, n_bins) -> a single sample
probs = torch.softmax(logits.float(), dim=-1)
bin_centers = _bin_centers(0.0, 1.0, probs.shape[-1]).to(device=probs.device)
return (probs * bin_centers).sum(dim=-1) * 100.0



def _select_best_sample(output: dict) -> int:
"""Select best diffusion sample by mean pLDDT."""
logits = output.get("plddt_logits")
Expand All @@ -185,7 +221,7 @@ def _select_best_sample(output: dict) -> int:
# (B, S, N_atoms, 50) → compute mean pLDDT per sample
probs = torch.softmax(logits[0], dim=-1)
n_bins = probs.shape[-1]
bin_centers = torch.linspace(0, 1, n_bins)
bin_centers = _bin_centers(n_bins, 0.0, 1.0)
plddt_per_atom = (probs * bin_centers).sum(dim=-1) # (S, N_atoms)
mean_plddt = plddt_per_atom.mean(dim=-1) # (S,)
return int(mean_plddt.argmax())
Comment on lines 214 to 227

@ducta3141 ducta3141 Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

def _select_best_sample(plddt_per_atom: torch.Tensor | None) -> int:
    """Select best diffusion sample by mean pLDDT.

    Reduces on the device and brings back only the winning index.
    """
    if plddt_per_atom is None:
        return 0
    return int(plddt_per_atom.mean(dim=-1).argmax())

Expand All @@ -206,7 +242,7 @@ def _compute_plddt(
logits = logits[0]
probs = torch.softmax(logits, dim=-1)
n_bins = probs.shape[-1]
bin_centers = torch.linspace(0, 1, n_bins)
bin_centers = _bin_centers(n_bins, 0.0, 1.0) # 50 bins -> 0.01, 0.03, ..., 0.99
plddt_per_atom = (probs * bin_centers).sum(dim=-1).numpy() * 100.0

# Aggregate to per-token
Expand All @@ -221,7 +257,12 @@ def _compute_plddt(
return plddt / counts


def _compute_ptm(output: dict, best_idx: int, n_tokens: int) -> float:
def _compute_ptm(
output: dict,
best_idx: int,
n_tokens: int,
has_frame: torch.Tensor | None = None,
) -> float:
"""Compute predicted TM-score from PAE logits."""
logits = output.get("pae_logits")
if logits is None:
Expand All @@ -232,10 +273,16 @@ def _compute_ptm(output: dict, best_idx: int, n_tokens: int) -> float:
elif logits.dim() == 4:
logits = logits[0]
logits = logits[:n_tokens, :n_tokens]
return float(_tm_score_from_pae_logits(logits, n_tokens))
return float(_tm_score_from_pae_logits(logits, n_tokens, has_frame=has_frame))
Comment on lines 260 to +276

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
def _compute_ptm(logits: torch.Tensor | None, n_tokens: int, has_frame: torch.Tensor | None = None) -> float:
"""Compute predicted TM-score from PAE logits."""
if logits is None:
return float("nan")
return _tm_score_from_pae_logits(logits, n_tokens, has_frame=has_frame)



def _compute_iptm(output: dict, best_idx: int, n_tokens: int, chain_indices: np.ndarray) -> float:
def _compute_iptm(
output: dict,
best_idx: int,
n_tokens: int,
chain_indices: np.ndarray,
has_frame: torch.Tensor | None = None,
) -> float:
"""Compute interface pTM from PAE logits (inter-chain pairs only)."""
logits = output.get("pae_logits")
if logits is None:
Expand All @@ -256,30 +303,66 @@ def _compute_iptm(output: dict, best_idx: int, n_tokens: int, chain_indices: np.
if inter_mask.sum() == 0:
return float("nan")

return float(_tm_score_from_pae_logits(logits, n_tokens, mask=inter_mask))
return float(_tm_score_from_pae_logits(logits, n_tokens, mask=inter_mask, has_frame=has_frame))
Comment on lines 279 to +306

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
def _compute_iptm(
logits: torch.Tensor | None,
n_tokens: int,
chain_indices: np.ndarray,
has_frame: torch.Tensor | None = None,
) -> float:
"""Compute interface pTM from PAE logits (inter-chain pairs only)."""
if logits is None:
return float("nan")
# Score only pairs that cross a chain boundary; a single chain has no interface.
ci = torch.as_tensor(chain_indices, dtype=torch.long, device=logits.device)
pair_mask = (ci.unsqueeze(-1) != ci.unsqueeze(-2)).to(dtype=torch.float32)
if not bool(pair_mask.any()):
return float("nan")
return _tm_score_from_pae_logits(logits, n_tokens, pair_mask=pair_mask, has_frame=has_frame)



def _tm_score_from_pae_logits(
logits: torch.Tensor,
n_tokens: int,
mask: torch.Tensor = None,
mask: torch.Tensor | None = None,
has_frame: torch.Tensor | None = None,
) -> torch.Tensor:
"""Compute TM-score from PAE logits using AF2/AF3 formula."""
probs = torch.softmax(logits, dim=-1)
"""Compute pTM / ipTM from PAE logits (AF3 SI §5.9.1, eqs. 17-18).

For every aligned token ``i`` the expected pairwise TM term
``E[1 / (1 + (e_ij / d0)^2)]`` is averaged over the scored tokens ``j``
(all tokens for pTM; tokens of other chains for ipTM, selected through
``mask[i, j]``), and the score is the *maximum* of these per-aligned-token
means over ``i`` -- the reduction used by AlphaFold2/3 and by upstream
OpenFold3 (``openfold3/core/metrics/confidence.py::compute_ptm``). A mean
over ``i`` (or over all pairs) is a lower bound of that value and is not
comparable with AF3-calibrated pTM / ipTM thresholds.

Args:
logits: (N, N, n_bins) PAE logits; row ``i`` is the aligned token.
n_tokens: number of tokens N used for ``d0`` (full complex).
mask: optional (N, N) 0/1 mask of scored pairs; ``None`` scores all pairs (pTM).
has_frame: optional (N,) bool mask restricting the max over ``i`` to
tokens with a valid frame (upstream derives it from the predicted
coordinates); ``None`` treats every token as a valid aligned token.

Returns NaN when no aligned token is eligible (``has_frame`` given with no True
entry, e.g. a ligand-only query under the interim mask, or no scored pair):
there is no frame-eligible token to align on. Upstream OpenFold3 returns 0.0 in
that case (``masked_fill`` then ``max``) and Protenix returns zeros; NaN is used
here so that ``FoldingOutput.get_scores()`` reports ``None`` rather than a
misleading 0, as it already does for the ipTM of single-chain inputs.
"""
probs = torch.softmax(logits.float(), dim=-1)
n_bins = probs.shape[-1]
bin_centers = torch.linspace(0, 32, n_bins) # 64 bins, 0-32 Å
bin_centers = _bin_centers(n_bins, 0.0, 32.0) # 64 bins on [0, 32] Å -> 0.25 ... 31.75

# d0 = 1.24 * (max(N, 19) - 15)^(1/3) - 1.8
d0 = 1.24 * (max(n_tokens, 19) - 15) ** (1.0 / 3.0) - 1.8
d0 = max(d0, 0.01)

# TM-score per pair: 1 / (1 + (d/d0)^2)
# TM-score term per pair: E_bins[1 / (1 + (e/d0)^2)]
tm_per_bin = 1.0 / (1.0 + (bin_centers / d0) ** 2)
tm_per_pair = (probs * tm_per_bin).sum(dim=-1) # (N, N)

if mask is not None:
return (tm_per_pair * mask).sum() / mask.sum().clamp(min=1)
return tm_per_pair.mean()
if mask is None:
mask = torch.ones_like(tm_per_pair)
mask = mask.to(dtype=tm_per_pair.dtype)

# Mean over scored tokens j for each aligned token i, then max over i.
n_scored = mask.sum(dim=-1) # (N,)
tm_per_aligned = (tm_per_pair * mask).sum(dim=-1) / n_scored.clamp(min=1)
valid = n_scored > 0
if has_frame is not None:
valid = valid & has_frame.to(device=valid.device, dtype=torch.bool)
if not bool(valid.any()):
return torch.tensor(float("nan"))
return tm_per_aligned[valid].max()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
probs = torch.softmax(logits.float(), dim=-1)
n_bins = probs.shape[-1]
# d0 = 1.24 * (max(N, 19) - 15)^(1/3) - 1.8, so the N floor of 19 keeps d0 > 0
d0 = 1.24 * (max(n_tokens, 19) - 15) ** (1.0 / 3.0) - 1.8
# Expected TM term per pair: E_bins[1 / (1 + (e_ij / d0)^2)]
bin_centers = _bin_centers(0.0, 32.0, n_bins).to(device=probs.device)
tm_per_bin = 1.0 / (1.0 + (bin_centers / d0) ** 2)
tm_per_pair = (probs * tm_per_bin).sum(dim=-1) # (N, N)
# Mean over the scored tokens j, for each aligned token i
if mask is None:
mask = torch.ones_like(tm_per_pair)
n_scored = mask.sum(dim=-1) # (N,)
tm_per_aligned = (tm_per_pair * mask).sum(dim=-1) / n_scored.clamp(min=1)
# Maximum over the aligned tokens that are eligible and have something to score
eligible = n_scored > 0
if has_frame is not None:
eligible = eligible & has_frame.to(device=eligible.device)
if not bool(eligible.any()):
return float("nan")
return float(tm_per_aligned[eligible].max())


def _compute_pae(output: dict, best_idx: int, n_tokens: int) -> np.ndarray | None:
Expand All @@ -295,7 +378,7 @@ def _compute_pae(output: dict, best_idx: int, n_tokens: int) -> np.ndarray | Non
logits = logits[:n_tokens, :n_tokens]
probs = torch.softmax(logits, dim=-1)
n_bins = probs.shape[-1]
bin_centers = torch.linspace(0, 32, n_bins)
bin_centers = _bin_centers(n_bins, 0.0, 32.0) # 64 bins -> 0.25, 0.75, ..., 31.75 Å
pae = (probs * bin_centers).sum(dim=-1).numpy()
return np.round(pae, 3)
Comment on lines 368 to 383

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

def _compute_pae(logits: torch.Tensor | None) -> np.ndarray | None:
    """Compute PAE matrix from PAE logits.

    Reducing the bin axis before the host copy makes the transfer ``n_bins``
    times smaller: an (N_token, N_token) matrix rather than the logits.
    """
    if logits is None:
        return None
    probs = torch.softmax(logits.float(), dim=-1)
    n_bins = probs.shape[-1]
    bin_centers = _bin_centers(0.0, 32.0, n_bins).to(device=probs.device)
    pae = (probs * bin_centers).sum(dim=-1).cpu().numpy()
    return np.round(pae, 3)


Expand Down
Loading