Buckets:
| """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 | |
| 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.