rescore.py: resolve v5-S2 ckpts under checkpoints/v5s2/seed_*/best_phase2.pt (HF layout) with fallback to results/v5_training/10seed/... (dev layout)
Browse files- code/scripts/rescore.py +17 -5
code/scripts/rescore.py
CHANGED
|
@@ -51,9 +51,20 @@ from utils.pdb_utils import (
|
|
| 51 |
# 3-seed ensemble (paper: v5-S2 TS-S2, seeds 1024 / 5555 / 789)
|
| 52 |
# --------------------------------------------------------------------------
|
| 53 |
SEEDS = [1024, 5555, 789]
|
| 54 |
-
|
|
|
|
|
|
|
|
|
|
| 55 |
ESM_DIR = str(BASE / "data/esm2_embeddings")
|
| 56 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
# Per-target holo/apo canonical PDB pair (extend as needed).
|
| 58 |
TARGETS = {
|
| 59 |
"cam": {"holo": "data/pdbs/cam_holo/3CLN.pdb", "apo": "data/pdbs/cam_apo/1CFD.pdb", "chain": "A", "esm_target": "cam"},
|
|
@@ -197,9 +208,10 @@ def rescore(target, design_dir, gpu, out_path):
|
|
| 197 |
|
| 198 |
results = {d["id"]: {} for d in extracted}
|
| 199 |
for seed in SEEDS:
|
| 200 |
-
ckpt =
|
| 201 |
-
if
|
| 202 |
-
|
|
|
|
| 203 |
continue
|
| 204 |
logger.info(f"Scoring seed {seed}...")
|
| 205 |
dq = DifferentiableQTheta(checkpoint_path=ckpt, device=device, esm_dir=ESM_DIR)
|
|
@@ -237,7 +249,7 @@ def rescore(target, design_dir, gpu, out_path):
|
|
| 237 |
|
| 238 |
out = {
|
| 239 |
"target": target, "holo_pdb": holo_pdb, "apo_pdb": apo_pdb,
|
| 240 |
-
"seeds": SEEDS, "checkpoints": {str(s):
|
| 241 |
"n_designs": len(extracted), "per_design": results,
|
| 242 |
}
|
| 243 |
Path(out_path).parent.mkdir(parents=True, exist_ok=True)
|
|
|
|
| 51 |
# 3-seed ensemble (paper: v5-S2 TS-S2, seeds 1024 / 5555 / 789)
|
| 52 |
# --------------------------------------------------------------------------
|
| 53 |
SEEDS = [1024, 5555, 789]
|
| 54 |
+
CKPT_LAYOUTS = [
|
| 55 |
+
str(BASE / "checkpoints/v5s2/seed_{seed}/best_phase2.pt"), # HF release layout
|
| 56 |
+
str(BASE / "results/v5_training/10seed/seed_{seed}/best_phase2.pt"), # in-repo dev layout
|
| 57 |
+
]
|
| 58 |
ESM_DIR = str(BASE / "data/esm2_embeddings")
|
| 59 |
|
| 60 |
+
|
| 61 |
+
def _resolve_ckpt(seed):
|
| 62 |
+
for tmpl in CKPT_LAYOUTS:
|
| 63 |
+
p = tmpl.format(seed=seed)
|
| 64 |
+
if Path(p).exists():
|
| 65 |
+
return p
|
| 66 |
+
return None
|
| 67 |
+
|
| 68 |
# Per-target holo/apo canonical PDB pair (extend as needed).
|
| 69 |
TARGETS = {
|
| 70 |
"cam": {"holo": "data/pdbs/cam_holo/3CLN.pdb", "apo": "data/pdbs/cam_apo/1CFD.pdb", "chain": "A", "esm_target": "cam"},
|
|
|
|
| 208 |
|
| 209 |
results = {d["id"]: {} for d in extracted}
|
| 210 |
for seed in SEEDS:
|
| 211 |
+
ckpt = _resolve_ckpt(seed)
|
| 212 |
+
if ckpt is None:
|
| 213 |
+
tried = [t.format(seed=seed) for t in CKPT_LAYOUTS]
|
| 214 |
+
logger.warning(f"Missing checkpoint for seed {seed} in any of: {tried}; skip")
|
| 215 |
continue
|
| 216 |
logger.info(f"Scoring seed {seed}...")
|
| 217 |
dq = DifferentiableQTheta(checkpoint_path=ckpt, device=device, esm_dir=ESM_DIR)
|
|
|
|
| 249 |
|
| 250 |
out = {
|
| 251 |
"target": target, "holo_pdb": holo_pdb, "apo_pdb": apo_pdb,
|
| 252 |
+
"seeds": SEEDS, "checkpoints": {str(s): _resolve_ckpt(s) for s in SEEDS},
|
| 253 |
"n_designs": len(extracted), "per_design": results,
|
| 254 |
}
|
| 255 |
Path(out_path).parent.mkdir(parents=True, exist_ok=True)
|