senba / inference_example.py
Kalipatnapu
Senba release: mars-fm-e010-a07350-exactlag-128x32-20260825-s1-v1 (validation-only)
112fa51 verified
Raw History Blame Contribute Delete
8.68 kB
#!/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())