File size: 4,037 Bytes
6cc35b0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Predicted-aligned-error scores and training loss."""

from __future__ import annotations

import torch
import torch.nn.functional as F
from torch import Tensor

from .esmfold2_affine3d import Affine3D

_CPU_DEVICE = torch.device("cpu")


def _compute_pae_masks(mask: Tensor) -> Tensor:
    residue_mask = mask.bool()
    return residue_mask.unsqueeze(-1) & residue_mask.unsqueeze(-2)


def _pae_bins(
    max_bin: float = 31,
    num_bins: int = 64,
    device: torch.device = _CPU_DEVICE,
) -> Tensor:
    """Return the representative distance for each PAE probability bin."""

    boundaries = torch.linspace(0, max_bin, steps=num_bins - 1, device=device)
    width = max_bin / (num_bins - 2)
    centers = boundaries + width / 2
    overflow_center = centers[-1:] + width
    return torch.cat((centers, overflow_center))


def _masked_probabilities(logits: Tensor, pair_mask: Tensor) -> Tensor:
    masked_logits = logits.masked_fill(
        ~pair_mask.unsqueeze(-1),
        torch.finfo(logits.dtype).min,
    )
    return masked_logits.softmax(dim=-1)


def masked_mean(
    mask: Tensor,
    value: Tensor,
    dim: int | tuple[int, ...] | None = None,
    eps: float = 1e-10,
) -> Tensor:
    """Average values over true entries of a broadcast-compatible mask."""

    weights = mask.expand_as(value)
    weighted_sum = torch.sum(weights * value, dim=dim)
    weight_sum = torch.sum(weights, dim=dim)
    return weighted_sum / (weight_sum + eps)


def compute_predicted_aligned_error(
    logits: Tensor,
    aa_mask: Tensor,
    sequence_id: Tensor | None = None,
    max_bin: float = 31,
) -> Tensor:
    """Convert PAE logits ``X`` with shape (..., l, l, n) to distances."""

    del sequence_id
    pair_mask = _compute_pae_masks(aa_mask)
    probabilities = _masked_probabilities(logits, pair_mask)
    centers = _pae_bins(max_bin, logits.shape[-1], logits.device)
    return torch.sum(probabilities * centers, dim=-1)


@torch.no_grad()
def compute_tm(logits: Tensor, aa_mask: Tensor, max_bin: float = 31.0) -> Tensor:
    """Estimate TM score from pairwise PAE logits."""

    pair_mask = _compute_pae_masks(aa_mask)
    sequence_lengths = aa_mask.sum(dim=-1, keepdim=True)
    centers = _pae_bins(max_bin, logits.shape[-1], logits.device)
    distance_scale = 1.24 * (sequence_lengths.clamp_min(19) - 15) ** (1 / 3) - 1.8
    tm_weights = 1.0 / (1 + (centers / distance_scale.unsqueeze(-1)) ** 2)
    probabilities = _masked_probabilities(logits, pair_mask)
    score_per_pair = torch.sum(probabilities * tm_weights.unsqueeze(-2), dim=-1)
    score_per_anchor = masked_mean(pair_mask, score_per_pair, dim=-1)
    return score_per_anchor.max(dim=-1).values


def _local_coordinates(frames: Affine3D) -> Tensor:
    origins = frames.trans[..., None, :, :]
    return frames.invert()[..., None].apply(origins)


def tm_loss(
    logits: Tensor,
    pred_affine: Tensor,
    targ_affine: Tensor,
    targ_mask: Tensor,
    tm_mask: Tensor | None = None,
    sequence_id: Tensor | None = None,
    max_bin: float = 31,
) -> Tensor:
    """Cross-entropy loss for discretized aligned-position errors."""

    del sequence_id
    predicted_frames = Affine3D.from_tensor(pred_affine)
    target_frames = Affine3D.from_tensor(targ_affine)
    with torch.no_grad():
        squared_error = (
            (_local_coordinates(predicted_frames) - _local_coordinates(target_frames))
            .square()
            .sum(dim=-1)
        )
        boundaries = torch.linspace(
            0,
            max_bin,
            logits.shape[-1] - 1,
            device=logits.device,
        ).square()
        target_bins = (squared_error[..., None] > boundaries).sum(dim=-1).long()

    cross_entropy = F.cross_entropy(
        logits.movedim(3, 1),
        target_bins,
        reduction="none",
    )
    pair_mask = _compute_pae_masks(targ_mask)
    loss_per_sample = masked_mean(pair_mask, cross_entropy, dim=(-1, -2))
    if tm_mask is None:
        return loss_per_sample.mean()
    return masked_mean(tm_mask, loss_per_sample)