StapleBridge / scripts /train.py
pranamanam's picture Jingjie00's picture
Upload Staplebridge files (#1)
bb6d2aa
Raw
History Blame Contribute Delete
18.7 kB
#!/usr/bin/env python
"""Training entry point for the main StapleBridge model (Full Exact-SB).
This is the orchestrator that produced the released checkpoint:
* full 4020-lead training each epoch;
* full validation on every validation lead (111), every epoch;
* a fixed 10 epochs, no early stopping (all 10 always run);
* two checkpoints maintained independently, each updated after that epoch's
validation pass:
- ``best_kl.pt``: minimum ``q_star_vs_q_theta_kl`` -- the selection rule;
- ``best_pv.pt``: maximum ``mean_delta_penetrance_vs_original_lead``,
recorded for monitoring only and not used to select the released model;
* per-epoch logging of both metrics with their running bests.
``checkpoints/staplebridge_seed42_best.pt`` is the ``best_kl.pt`` of this run:
the epoch minimising ``q_star_vs_q_theta_kl`` on the validation split. Model,
loss, Exact-SB, property scoring, decoding and every other training setting are
read from the config.
"""
from __future__ import annotations
import argparse
import gc
import json
import random
import sys
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
import yaml
PACKAGE_ROOT = Path(__file__).resolve().parents[1]
if str(PACKAGE_ROOT) not in sys.path:
sys.path.insert(0, str(PACKAGE_ROOT))
from staplebridge.data.dataset import load_leads # noqa: E402
from staplebridge.hydrocarbon.exact_sb_cache import build_cache_from_config # noqa: E402
from staplebridge.hydrocarbon.plan_control import ( # noqa: E402
HydrocarbonPlanControlConfig, build_hydrocarbon_plan_head,
)
from staplebridge.hydrocarbon.property_energy import ( # noqa: E402
HydrocarbonPropertyEnergyConfig,
HydrocarbonPropertyScorer,
required_original_lead_properties,
)
from staplebridge.hydrocarbon.tokenizer import tokenize_sequence # noqa: E402
from staplebridge.training.main_loop import ( # noqa: E402
train_enabled_epoch, validate_enabled,
)
from staplebridge.training.records import write_json, write_jsonl # noqa: E402
from staplebridge.training.stack import ( # noqa: E402
build_energy, build_models, build_predictor, build_stack, load_config,
seed_everything, select_leads,
)
FULL_TRAIN_N = 4020
SEED = 42
EPOCHS = 10
KL_KEY = "q_star_vs_q_theta_kl"
PV_KEY = "mean_delta_penetrance_vs_original_lead"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", type=Path,
default=PACKAGE_ROOT / "configs/staplebridge_main.yaml")
parser.add_argument("--out-dir", type=Path,
default=PACKAGE_ROOT / "outputs/main_seed42")
parser.add_argument("--resume", type=Path, default=None,
help="Checkpoint to resume from (e.g. checkpoints/epoch_008.pt). "
"Restores model/optimizer/RNG so remaining epochs are identical "
"to an uninterrupted run; appends to the existing metrics/log.")
return parser.parse_args()
def require(condition: bool, message: str) -> None:
if not condition:
raise SystemExit(message)
def rng_payload() -> dict[str, Any]:
return {
"python_random_state": random.getstate(),
"numpy_random_state": np.random.get_state(),
"torch_rng_state": torch.get_rng_state(),
"cuda_rng_state_all": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None,
}
def build_original_cache(config: dict[str, Any], leads: list[Any], out_dir: Path) -> dict:
"""Compute the configured unedited-lead property cache once."""
path = out_dir / "original_linear_cache.jsonl"
wrapper, _ = build_predictor(config)
cache: dict[tuple[Any, ...], dict[str, Any]] = {}
scorer = HydrocarbonPropertyScorer(wrapper, original_linear_cache=cache)
property_cfg = HydrocarbonPropertyEnergyConfig.from_dict(
(((config.get("hydrocarbon") or {}).get("terminal_energy") or {}).get("property"))
)
properties = required_original_lead_properties(property_cfg)
started = time.perf_counter()
if (
property_cfg.enable_developability_constraints
or property_cfg.enable_halflife_preservation
or property_cfg.enable_joint_perm_halflife_support
):
scorer.prefetch(
properties,
[
scorer.original_linear_smiles(tokenize_sequence(lead.linear_sequence))
for lead in leads
],
)
for index, lead in enumerate(leads):
scorer.score_original_linear(
tokenize_sequence(lead.linear_sequence),
lead_key=str(lead.example_id),
properties=properties,
)
if (index + 1) % 512 == 0 or index + 1 == len(leads):
print(f"[original baseline] {index + 1}/{len(leads)}", flush=True)
rows = [{"lead_key": k[0], "tokens": list(k[1]), "scores": v} for k, v in cache.items()]
write_jsonl(path, rows)
del scorer, wrapper
import gc
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
print(f"[original baseline] built {len(cache)} entries in {time.perf_counter() - started:.1f}s", flush=True)
return cache
def load_original_cache(path: Path) -> dict[tuple[Any, ...], dict[str, Any]]:
cache: dict[tuple[Any, ...], dict[str, Any]] = {}
with path.open(encoding="utf-8") as handle:
for line in handle:
row = json.loads(line)
cache[(row["lead_key"], tuple(row["tokens"]))] = dict(row["scores"])
return cache
def save_checkpoint(path: Path, payload: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_suffix(path.suffix + ".tmp")
torch.save(payload, tmp)
tmp.replace(path)
def main() -> None:
args = parse_args()
config = load_config(args.config.resolve())
out_dir = args.out_dir.resolve()
resuming = args.resume is not None
if resuming:
require(out_dir.is_dir(), f"resume requires existing output dir: {out_dir}")
require(args.resume.is_file(), f"resume checkpoint not found: {args.resume}")
else:
# Never overwrite prior outputs: require a fresh/empty directory.
if out_dir.exists():
leftovers = [p for p in out_dir.iterdir()]
require(not leftovers, f"output directory not empty (refusing to overwrite): {out_dir}")
out_dir.mkdir(parents=True, exist_ok=True)
require(int(config.get("train_n", -1)) == FULL_TRAIN_N, "train_n must be 4020")
require(int(config["training"]["epochs"]) == EPOCHS, f"epochs must be {EPOCHS}")
device = torch.device(str(config["training"]["device"]))
if device.type == "cuda":
require(torch.cuda.is_available(), f"CUDA unavailable for {device}")
torch.cuda.set_device(device)
(out_dir / "resolved_config.yaml").write_text(yaml.safe_dump(config, sort_keys=False))
seed_everything(SEED)
# ---- data: full train (4020) + FULL valid (all 111, not the 32 subset) --
train_leads = select_leads(
Path(config["data"]["root"]) / config["data"]["train_file"],
FULL_TRAIN_N, int(config["data"]["max_lead_length"]),
)
valid_cap = int(config["validation"]["max_lead_length"])
valid_leads = [
lead for lead in load_leads(Path(config["data"]["root"]) / config["data"]["valid_file"])
if len(lead.linear_sequence) <= valid_cap
]
print(f"[data] train={len(train_leads)} full_valid={len(valid_leads)} (cap len<={valid_cap})", flush=True)
# ---- original-lead penetrance cache (train + valid) for delta-PV --------
original_path = out_dir / "original_linear_cache.jsonl"
if resuming and original_path.is_file():
original_cache = load_original_cache(original_path)
print(f"[original baseline] reused {len(original_cache)} entries from {original_path}", flush=True)
else:
original_cache = build_original_cache(config, train_leads + valid_leads, out_dir)
# ---- model + energy stack (identical construction to the standard run) --
stack = build_stack(config, SEED)
wrapper, _ = build_predictor(config)
scorer = HydrocarbonPropertyScorer(wrapper, original_linear_cache=original_cache)
energy_fn = build_energy(config, stack, scorer)
policy, value, kernel, optimizer = build_models(config, stack, device)
head = build_hydrocarbon_plan_head(
config, int(config["model"]["emb_dim"]), device,
esm2_prior=stack["reference_priors"].peptide,
)
optimizer.add_param_group({"params": list(head.parameters())})
parameters = list(policy.parameters()) + list(value.parameters()) + list(head.parameters())
plan_cfg = HydrocarbonPlanControlConfig.from_config(config)
plan_rng = random.Random(SEED)
# ---- reuse the existing persistent q* cache (fingerprint unchanged) -----
exact_sb_cache = build_cache_from_config(config, catalog=stack["catalog"], repo_root=PACKAGE_ROOT)
if exact_sb_cache.enabled:
print(f"[exact-sb cache] {json.dumps(exact_sb_cache.describe(), ensure_ascii=False)}", flush=True)
def make_payload(epoch: int) -> dict[str, Any]:
return {
"epoch": epoch, "config": config, "plan_control_enabled": True,
"policy_state_dict": policy.state_dict(), "value_state_dict": value.state_dict(),
"plan_head_state_dict": head.state_dict(), "optimizer_state_dict": optimizer.state_dict(),
"plan_rng_state": plan_rng.getstate(), **rng_payload(),
}
metrics_path = out_dir / "metrics.jsonl"
history: list[dict[str, Any]] = []
best_kl = {"value": float("inf"), "epoch": None}
best_pv = {"value": float("-inf"), "epoch": None}
start_epoch = 0
if resuming:
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
policy.load_state_dict(ckpt["policy_state_dict"])
value.load_state_dict(ckpt["value_state_dict"])
head.load_state_dict(ckpt["plan_head_state_dict"])
optimizer.load_state_dict(ckpt["optimizer_state_dict"])
plan_rng.setstate(ckpt["plan_rng_state"])
random.setstate(ckpt["python_random_state"])
np.random.set_state(ckpt["numpy_random_state"])
torch.set_rng_state(ckpt["torch_rng_state"].cpu())
if torch.cuda.is_available() and ckpt.get("cuda_rng_state_all") is not None:
torch.cuda.set_rng_state_all([s.cpu() for s in ckpt["cuda_rng_state_all"]])
start_epoch = int(ckpt["epoch"])
# Rebuild running bests + history from the persisted per-epoch metrics so
# best_kl/best_pv provenance carries across the resume boundary.
for line in metrics_path.read_text().splitlines():
rec = json.loads(line)
if int(rec["epoch"]) <= start_epoch:
history.append(rec)
if history:
last = history[-1]
best_kl = {"value": float(last["best_kl_value"]), "epoch": int(last["best_kl_epoch"])}
best_pv = {"value": float(last["best_pv_value"]), "epoch": int(last["best_pv_epoch"])}
print(
f"[resume] from {args.resume} completed_epoch={start_epoch}; "
f"best_kl={best_kl['value']:.6f}@ep{best_kl['epoch']} "
f"best_pv={best_pv['value']:.6f}@ep{best_pv['epoch']}",
flush=True,
)
run_started = time.perf_counter()
for epoch in range(start_epoch, EPOCHS):
policy.train(); value.train(); head.train()
train_started = time.perf_counter()
rows, train_metrics = train_enabled_epoch(
train_leads, config, stack, energy_fn, policy, kernel, optimizer,
parameters, head, plan_cfg, plan_rng, epoch, exact_sb_cache=exact_sb_cache,
)
train_seconds = time.perf_counter() - train_started
write_jsonl(out_dir / "training" / f"epoch_{epoch + 1:03d}_candidates.jsonl", rows)
joint_train_audit = train_metrics.get("joint_perm_halflife_support")
if joint_train_audit and epoch == 0:
write_jsonl(
out_dir / "joint_support_audit" / "train4020_per_lead.jsonl",
joint_train_audit["per_lead"],
)
write_json(
out_dir / "joint_support_audit" / "train4020_summary.json",
{k: v for k, v in joint_train_audit.items() if k != "per_lead"},
)
policy.eval(); value.eval(); head.eval()
valid_started = time.perf_counter()
selected, validation = validate_enabled(
valid_leads, config, stack, energy_fn, policy, kernel, head,
exact_sb_cache=exact_sb_cache,
)
valid_seconds = time.perf_counter() - valid_started
valid_dir = out_dir / "validation" / f"epoch_{epoch + 1:03d}"
write_jsonl(valid_dir / "selected.jsonl", selected)
write_json(valid_dir / "summary.json", validation)
joint_valid_audit = validation.get("joint_perm_halflife_support")
if joint_valid_audit:
write_jsonl(
valid_dir / "joint_support_per_lead.jsonl",
joint_valid_audit["per_lead"],
)
write_json(
valid_dir / "joint_support_summary.json",
{k: v for k, v in joint_valid_audit.items() if k != "per_lead"},
)
kl = float(validation[KL_KEY])
pv = float(validation[PV_KEY])
guard_cfg = dict(config.get("guardrails") or {})
both_topologies_present = bool(
validation.get("s5_s5_i4_count", 0) > 0
and validation.get("r8_s5_i7_count", 0) > 0
)
checkpoint_eligible = bool(
not guard_cfg.get("require_both_topologies", False)
or both_topologies_present
)
improved_kl = checkpoint_eligible and kl < best_kl["value"]
improved_pv = checkpoint_eligible and pv > best_pv["value"]
if improved_kl:
best_kl = {"value": kl, "epoch": epoch + 1}
save_checkpoint(out_dir / "checkpoints" / "best_kl.pt", make_payload(epoch + 1))
if improved_pv:
best_pv = {"value": pv, "epoch": epoch + 1}
save_checkpoint(out_dir / "checkpoints" / "best_pv.pt", make_payload(epoch + 1))
# Per-epoch + latest snapshots so no epoch is lost (fresh dir; nothing
# is overwritten across runs).
payload = make_payload(epoch + 1)
save_checkpoint(out_dir / "checkpoints" / f"epoch_{epoch + 1:03d}.pt", payload)
save_checkpoint(out_dir / "checkpoints" / "latest.pt", payload)
record = {
"epoch": epoch + 1,
"valid_kl": kl,
"valid_delta_penetrance": pv,
"best_kl_epoch": best_kl["epoch"], "best_kl_value": best_kl["value"], "kl_improved": improved_kl,
"best_pv_epoch": best_pv["epoch"], "best_pv_value": best_pv["value"], "pv_improved": improved_pv,
"valid_product_penetrance": validation.get("mean_product_penetrance"),
"valid_top1_chemistry_valid_rate": validation.get("top1_chemistry_valid_rate"),
"valid_top1_stapled_rate": validation.get("top1_stapled_rate"),
"valid_q_star_top1_agreement": validation.get("q_star_top1_agreement"),
"valid_q_star_spearman": validation.get("q_star_spearman"),
"valid_edit_distance": validation.get("mean_weighted_edit_distance"),
"both_topologies_present": both_topologies_present,
"checkpoint_eligible": checkpoint_eligible,
"train_loss": train_metrics.get("loss"),
"train_plan_loss": train_metrics.get("plan_loss"),
"train_q_star_vs_q_theta_kl": train_metrics.get("q_star_vs_q_theta_kl"),
"train_seconds": train_seconds,
"valid_seconds": valid_seconds,
"epoch_seconds": train_seconds + valid_seconds,
"train_stage_seconds": train_metrics.get("stage_seconds"),
"exact_sb_cache": exact_sb_cache.describe() if exact_sb_cache.enabled else {"enabled": False},
"train_joint_perm_halflife_support": (
{k: v for k, v in (joint_train_audit or {}).items() if k != "per_lead"}
if joint_train_audit
else None
),
"valid_joint_perm_halflife_support": (
{
k: v
for k, v in (joint_valid_audit or {}).items()
if k != "per_lead"
}
if joint_valid_audit
else None
),
}
history.append(record)
write_jsonl(metrics_path, [record], mode="a")
print(
f"[epoch {epoch + 1}/{EPOCHS}] "
f"valid_KL={kl:.6f} (best {best_kl['value']:.6f} @ep{best_kl['epoch']}"
f"{' NEW' if improved_kl else ''}) "
f"valid_deltaPV={pv:.6f} (best {best_pv['value']:.6f} @ep{best_pv['epoch']}"
f"{' NEW' if improved_pv else ''}) "
f"epoch_seconds={record['epoch_seconds']:.1f}",
flush=True,
)
# Semantics-neutral memory hygiene: reclaim inter-epoch CUDA cache /
# Python garbage so allocator fragmentation does not accumulate across
# the 10 epochs. Does not touch weights, RNG, or any cached value.
gc.collect()
if device.type == "cuda":
torch.cuda.empty_cache()
print(
f"[gpu] epoch {epoch + 1} allocated={torch.cuda.memory_allocated(device) / 2**20:.0f}MiB "
f"reserved={torch.cuda.memory_reserved(device) / 2**20:.0f}MiB",
flush=True,
)
summary = {
"exp_name": config.get("exp_name"),
"epochs": EPOCHS,
"early_stopping": False,
"train_n": FULL_TRAIN_N,
"n_valid_leads": len(valid_leads),
"seed": SEED,
"best_kl": best_kl,
"best_pv": best_pv,
"kl_curve": [(r["epoch"], r["valid_kl"]) for r in history],
"pv_curve": [(r["epoch"], r["valid_delta_penetrance"]) for r in history],
"runtime_seconds": time.perf_counter() - run_started,
"gpu_peak_mib": (torch.cuda.max_memory_allocated(device) / 2 ** 20) if device.type == "cuda" else 0,
"checkpoints": {
"best_kl": str(out_dir / "checkpoints" / "best_kl.pt"),
"best_pv": str(out_dir / "checkpoints" / "best_pv.pt"),
},
"history": history,
}
write_json(out_dir / "run_summary.json", summary)
print("\n" + json.dumps({k: v for k, v in summary.items() if k != "history"}, indent=2, default=str), flush=True)
if __name__ == "__main__":
main()