dmitry-rov/apt-repro / scripts /claim3_designability.py
dmitry-rov's picture
download
raw
6.22 kB
"""Claim 3 — designability of APT-generated backbones.
Pipeline (standard self-consistency designability):
1. Generate N CA backbones with the APT LM + tokenizer.
sampling = minp (released generate.py); decoding uses classifier annealing
(cfg_fn = 1 - t, the paper's mechanism) via decode_with_cfg_fn.
2. ProteinMPNN (CA-only) designs K sequences per backbone.
3. ESMFold refolds each sequence.
4. scRMSD = min over sequences of CA-RMSD(ESMFold, generated backbone).
designable if scRMSD < 2 Å. designability = fraction designable.
Note: released generate.py has no *entropy* sampler (paper's exact Table-4 method),
so this measures designability of APT generations under the available sampler +
classifier annealing, not the exact 0.871 configuration.
"""
from __future__ import annotations
import json
import os
import subprocess
import sys
from pathlib import Path
import numpy as np
import torch
from apt.models import APTLanguageModel, APTTokenizer
from apt.generate import generate_spda
from apt.utils import kabsch_rmsd
CKPT = Path(os.environ.get("CKPT_DIR", "/tmp/apt_ckpts"))
TOK = CKPT / "tokenizer128.pt"
LM = CKPT / "lm128.pt"
OUT = Path(os.environ.get("OUT_DIR", "/tmp/out"))
N_BACKBONES = int(os.environ.get("N_BACKBONES", "50"))
K_SEQS = int(os.environ.get("K_SEQS", "8"))
MPNN_DIR = Path("/tmp/ProteinMPNN")
def write_ca_pdb(path: Path, coords: np.ndarray) -> None:
lines = []
for i, (x, y, z) in enumerate(coords, start=1):
lines.append(
f"ATOM {i:5d} CA ALA A{i:4d} {x:8.3f}{y:8.3f}{z:8.3f} 1.00 20.00 C"
)
lines.append("END")
path.write_text("\n".join(lines) + "\n")
def setup_mpnn() -> None:
if not MPNN_DIR.exists():
subprocess.run(
["git", "clone", "--depth", "1", "https://github.com/dauparas/ProteinMPNN", str(MPNN_DIR)],
check=True,
)
def mpnn_design(pdb: Path, out_dir: Path, k: int) -> list[str]:
out_dir.mkdir(parents=True, exist_ok=True)
subprocess.run(
[sys.executable, str(MPNN_DIR / "protein_mpnn_run.py"),
"--ca_only", "--pdb_path", str(pdb), "--out_folder", str(out_dir),
"--num_seq_per_target", str(k), "--sampling_temp", "0.1",
"--batch_size", "1", "--seed", "37"],
check=True, capture_output=True,
)
fa = out_dir / "seqs" / (pdb.stem + ".fa")
seqs = []
for line in fa.read_text().splitlines():
if line and not line.startswith(">"):
seqs.append(line.strip())
# first record is the native (poly-ALA) sequence; drop it
return seqs[1:] if len(seqs) > 1 else seqs
def load_esmfold(dev):
from transformers import AutoTokenizer, EsmForProteinFolding
tok = AutoTokenizer.from_pretrained("facebook/esmfold_v1")
model = EsmForProteinFolding.from_pretrained("facebook/esmfold_v1")
model = model.to(dev).eval()
model.esm = model.esm.half()
model.trunk.set_chunk_size(64)
return tok, model
@torch.no_grad()
def esmfold_ca(tok, model, seq: str, dev) -> np.ndarray:
ids = tok([seq], return_tensors="pt", add_special_tokens=False)["input_ids"].to(dev)
out = model(ids)
pos = out["positions"][-1, 0] # (L, 14, 3)
return pos[:, 1, :].float().cpu().numpy() # CA
def main() -> None:
dev = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(37)
setup_mpnn()
tokenizer = APTTokenizer.from_pretrained(TOK).to(dev).eval()
model = APTLanguageModel.from_pretrained(LM, TOK).to(dev).eval()
# 1. generate backbones (one at a time; lengths vary)
print(f"generating {N_BACKBONES} backbones...")
backbones = []
tmp = Path("/tmp/gen"); tmp.mkdir(exist_ok=True)
while len(backbones) < N_BACKBONES:
seqs = generate_spda(model, batch_size=1, sampling_method="minp",
threshold=0.15, temperature=1.0)
idx_BL = seqs[0].unsqueeze(0)
L = idx_BL.size(1)
if L < 20 or L > 128:
continue
with torch.no_grad():
coords = tokenizer.decode_with_cfg_fn(
idx_BL, cfg_fn=lambda t: 1 - t, true_length=L,
n_steps=100, noise_weight=0.1, score_weight=1.0)
ca = coords.squeeze(0).cpu().numpy() * 10.0
backbones.append(ca)
if len(backbones) % 10 == 0:
print(f" {len(backbones)}/{N_BACKBONES}")
# 2-4. design + fold + score
print("loading ESMFold...")
eft, efm = load_esmfold(dev)
results = []
for i, ca in enumerate(backbones):
pdb = tmp / f"bb_{i}.pdb"
write_ca_pdb(pdb, ca)
try:
seqs = mpnn_design(pdb, tmp / f"mpnn_{i}", K_SEQS)
except subprocess.CalledProcessError as e:
print(f" mpnn fail bb{i}: {e.stderr.decode()[:200]}")
continue
bb_t = torch.from_numpy(ca).float().unsqueeze(0) / 10.0
best = 99.0
for s in seqs:
try:
pred = esmfold_ca(eft, efm, s, dev)
if pred.shape[0] != ca.shape[0]:
continue
pred_t = torch.from_numpy(pred).float().unsqueeze(0) / 10.0
r = float(kabsch_rmsd(pred_t, bb_t).item()) * 10.0
best = min(best, r)
except Exception as e: # noqa: BLE001
print(f" esmfold fail bb{i}: {e}")
results.append({"i": i, "L": int(ca.shape[0]), "scRMSD": round(best, 3),
"designable": best < 2.0})
print(f" bb{i} L={ca.shape[0]} scRMSD={best:.3f} designable={best < 2.0}")
scr = np.array([r["scRMSD"] for r in results])
des = np.array([r["designable"] for r in results])
summary = {
"n_backbones": len(results), "k_seqs": K_SEQS,
"designability": round(float(des.mean()), 4),
"scRMSD_mean": round(float(scr.mean()), 4),
"scRMSD_median": round(float(np.median(scr)), 4),
"sampler": "minp+classifier_annealing(cfg=1-t)",
}
print("RESULT_JSON " + json.dumps(summary))
OUT.mkdir(parents=True, exist_ok=True)
(OUT / "claim3_designability.json").write_text(json.dumps({"summary": summary, "per": results}, indent=2))
if __name__ == "__main__":
main()

Xet Storage Details

Size:
6.22 kB
·
Xet hash:
44e1f9cd7a27bb98b7206197cee6d5690e447b184cd41224d63b16c64816a337

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.