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())
|