Kalipatnapu
Senba release: mars-fm-e010-a07350-exactlag-128x32-20260825-s1-v1 (validation-only)
112fa51 verified Download inference_example.py from vkali08/senba: direct link, hf CLI and curl.
- Browser
- Download file 8.68 kB
-
https://huggingface.co/vkali08/senba/resolve/main/inference_example.py
- Command line
-
hf download hf://vkali08/senba/inference_example.py
-
curl -L -o inference_example.py https://huggingface.co/vkali08/senba/resolve/main/inference_example.py
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()) | |