File size: 8,683 Bytes
112fa51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python
"""Sample Senba (adapted MarS-FM) transition endpoints for one protein structure.

Senba is a MarS-FM checkpoint, so it loads with the pinned MarS-FM source
(https://github.com/valence-labs/mars-fm at 4fc17d86cd7e22dda6353f2f1f6bf4b40c9ce946).
Put that repository on PYTHONPATH, then run:

    python inference_example.py --pdb 1ubq.pdb --chain A \
        --checkpoint senba.ckpt \
        --expect-sha256 f02dcf6a97848c834e3c0be7d848c1068d1d0eb7976d8d6bcfbc98d4114e9c7c \
        --samples 8 --ode-steps 10 --seed 20260904 --out senba-1ubq.pdb

What you get: a multi-model PDB whose MODEL 1 is the input heavy-atom structure
and whose MODELS 2..K+1 are K independent samples of where the model places the
chain after the trained lag (50 mdCATH frames at 450 K). The samples are not a
trajectory, carry no physical clock, and were validated only on single-chain
mdCATH domains of at most 256 residues. Every claim about Senba is bounded to
that validation evidence; see README.md.

Input requirements: one protein chain, standard amino acids only, every heavy
atom present (crystal structures with truncated side chains are rejected rather
than silently zero-filled). HETATM records, hydrogens, alternate locations other
than A, and other chains are dropped before parsing.
"""

from __future__ import annotations

import argparse
import hashlib
import sys
from pathlib import Path

import numpy as np
import torch

torch.serialization.add_safe_globals([argparse.Namespace])

from mars.data.geometry import atom14_to_atom37, atom14_to_frames, atom37_to_torsions  # noqa: E402
from mars.model.module import MarSModule  # noqa: E402
from mars.utils import atom14_to_pdb  # noqa: E402
from mars.vendored.openfold import protein  # noqa: E402
from mars.vendored.openfold import residue_constants as rc  # noqa: E402

VALIDATED_MAX_RESIDUES = 256
TRAINED_LAG_FRAMES = 50
TRAINED_TEMPERATURE_K = 450


def sha256_file(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        while chunk := handle.read(8 * 1024 * 1024):
            digest.update(chunk)
    return digest.hexdigest()


def clean_pdb_text(pdb_text: str, chain_id: str | None) -> str:
    """Keep ATOM records of one chain with primary alternate locations only."""
    kept: list[str] = []
    for line in pdb_text.splitlines():
        if not line.startswith("ATOM  "):
            continue
        if len(line) < 54:
            raise ValueError("truncated ATOM record in input PDB")
        if line[16] not in (" ", "A"):
            continue
        if chain_id is not None and line[21].strip() != chain_id:
            continue
        kept.append(line[:16] + " " + line[17:])
    if not kept:
        raise ValueError("no ATOM records matched the requested chain")
    return "\n".join(kept) + "\nEND\n"


def pdb_to_atom14(pdb_text: str, chain_id: str | None) -> tuple[np.ndarray, str]:
    prot = protein.from_pdb_string(clean_pdb_text(pdb_text, chain_id))
    if prot.chain_index is not None and len(np.unique(prot.chain_index)) != 1:
        raise ValueError("the structure has several chains; pass --chain")
    if prot.aatype.size == 0:
        raise ValueError("no residues were parsed")
    if np.any(prot.aatype >= rc.restype_num):
        raise ValueError("non-standard residues are not supported by MarS-FM")
    n_res = int(prot.aatype.shape[0])
    atom14 = np.zeros((n_res, 14, 3), dtype=np.float32)
    for index, aatype in enumerate(prot.aatype.tolist()):
        names = rc.restype_name_to_atom14_names[rc.restype_1to3[rc.restypes[aatype]]]
        for slot, name in enumerate(names):
            if not name:
                continue
            atom_index = rc.atom_order[name]
            if prot.atom_mask[index, atom_index] <= 0:
                raise ValueError(
                    f"residue {index + 1} ({rc.restypes[aatype]}) is missing heavy atom {name}; "
                    "supply a structure with complete side chains"
                )
            atom14[index, slot] = prot.atom_positions[index, atom_index]
    sequence = "".join(rc.restypes[aatype] for aatype in prot.aatype.tolist())
    return atom14, sequence


def build_batch(atom14: np.ndarray, sequence: str, samples: int, device: torch.device) -> dict:
    """Mirror the batch layout used by the frozen Senba validation runs."""
    seqres = torch.tensor([rc.restype_order[residue] for residue in sequence])
    seqres = seqres.unsqueeze(0).repeat(samples, 1)
    repeated = np.repeat(atom14[None], samples, axis=0).astype(np.float32, copy=False)
    frames = atom14_to_frames(torch.from_numpy(repeated))
    atom37 = torch.from_numpy(atom14_to_atom37(repeated, seqres)).float()
    torsions, torsion_mask = atom37_to_torsions(atom37, seqres)
    batch = {
        "torsions": torsions.unsqueeze(1),
        "torsion_mask": torsion_mask,
        "trans": frames._trans.unsqueeze(1),
        "rots": frames._rots._rot_mats.unsqueeze(1),
        "seqres": seqres,
        "mask": torch.ones_like(seqres, dtype=torch.float32),
    }
    return {key: value.to(device) for key, value in batch.items()}


def kabsch_rmsd(reference: np.ndarray, mobile: np.ndarray) -> float:
    """C-alpha RMSD after optimal superposition (no reflection)."""
    ref = reference - reference.mean(axis=0)
    mob = mobile - mobile.mean(axis=0)
    covariance = mob.T @ ref
    u, _, vt = np.linalg.svd(covariance)
    sign = np.sign(np.linalg.det(u @ vt))
    correction = np.diag([1.0, 1.0, sign])
    rotation = u @ correction @ vt
    aligned = mob @ rotation
    return float(np.sqrt(np.mean(np.sum((aligned - ref) ** 2, axis=1))))


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("--pdb", required=True, type=Path, help="input PDB file (heavy atoms, one chain)")
    parser.add_argument("--chain", default=None, help="chain identifier to use when the file has several")
    parser.add_argument("--checkpoint", required=True, type=Path, help="path to senba.ckpt")
    parser.add_argument("--expect-sha256", default=None, help="refuse to run unless the checkpoint hash matches")
    parser.add_argument("--samples", type=int, default=8, help="independent endpoint draws")
    parser.add_argument("--ode-steps", type=int, default=10, help="flow ODE steps (validation used 10)")
    parser.add_argument("--seed", type=int, default=0)
    parser.add_argument("--out", required=True, type=Path, help="multi-model PDB to write")
    parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
    args = parser.parse_args()

    if args.samples < 1 or args.ode_steps < 1:
        parser.error("--samples and --ode-steps must be positive")
    digest = sha256_file(args.checkpoint)
    if args.expect_sha256 is not None and digest != args.expect_sha256.lower().removeprefix("sha256:"):
        print(f"checkpoint hash {digest} does not match --expect-sha256", file=sys.stderr)
        return 2

    atom14, sequence = pdb_to_atom14(args.pdb.read_text(), args.chain)
    if len(sequence) > VALIDATED_MAX_RESIDUES:
        print(
            f"warning: {len(sequence)} residues exceeds the validated scope of "
            f"{VALIDATED_MAX_RESIDUES}; outputs are outside the published evidence",
            file=sys.stderr,
        )

    device = torch.device(args.device)
    torch.manual_seed(args.seed)
    if device.type == "cuda":
        torch.cuda.manual_seed_all(args.seed)
    model = MarSModule.load_from_checkpoint(str(args.checkpoint), map_location="cpu")
    model.eval().to(device)
    batch = build_batch(atom14, sequence, args.samples, device)
    with torch.no_grad():
        generated_atom14, _ = model.inference(batch, num_steps=args.ode_steps)
    samples = generated_atom14[:, 0].float().cpu().numpy()  # [K, L, 14, 3]

    aatype = np.array([rc.restype_order[residue] for residue in sequence])
    atom14_to_pdb(np.concatenate([atom14[None], samples], axis=0), aatype, str(args.out))

    reference_ca = atom14[:, 1]
    rmsds = [kabsch_rmsd(reference_ca, sample[:, 1]) for sample in samples]
    print(f"checkpoint sha256 {digest}")
    print(f"sequence length {len(sequence)}; lag {TRAINED_LAG_FRAMES} frames at {TRAINED_TEMPERATURE_K} K")
    print(f"wrote {args.out} with MODEL 1 = input and MODELS 2..{args.samples + 1} = independent draws")
    print("aligned C-alpha RMSD of each draw from the input (A): " + ", ".join(f"{value:.3f}" for value in rmsds))
    print("these draws are independent samples at one lag, not a trajectory or a physical-time prediction")
    return 0


if __name__ == "__main__":
    sys.exit(main())