StapleBridge / staplebridge /hydrocarbon /exact_sb_cache.py
pranamanam's picture Jingjie00's picture
Upload Staplebridge files (#1)
bb6d2aa
Raw History Blame Contribute Delete
28.6 kB
"""Persistent cache for the deterministic part of the Exact-SB plan target.
For a fixed lead and a fixed configuration, everything on the right-hand side of
q*(p|x) ∝ q_ref(p|x) · exp[-beta · E_T(x, p)]
is deterministic: the legal-plan enumeration, ``q_ref``, the PeptiVerse-backed
terminal energies, and therefore ``log q*`` itself. Only ``q_theta`` changes as
the plan head trains. This module persists the deterministic half so each epoch
recomputes just ``q_theta`` and then ``KL(q*||q_theta)``.
Deliberately **not** cached:
* ``q_theta`` — it is a function of the live model weights. Caching or freezing
it would silently stop training the plan head. Nothing in this module reads,
writes, or accepts ``q_theta``.
* Anything used for candidate ranking, decoding, or the legacy validation
metrics. The cache only replays plan order, ``reference_logp``, terminal
energies and ``log q*``; every consumer recomputes ``q_theta`` itself.
Correctness rests on the fingerprint: any change to the lead, catalog,
empirical prior, ``exact_sb_beta``, terminal-energy/property configuration,
PeptiVerse model set, or geometry/edit settings produces a different
fingerprint, so a stale entry is never served. Entries also store the plan
signatures they were built from and are rejected if the live enumeration
disagrees, which catches drift the fingerprint alone would miss.
"""
from __future__ import annotations
import hashlib
import inspect
import json
import os
import sqlite3
import threading
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Iterable, Sequence
import torch
from staplebridge.hydrocarbon.curriculum import HydrocarbonStaplePlan
#: Bumped whenever the stored payload's meaning changes. Old rows then miss.
CACHE_SCHEMA_VERSION = 1
def _stable_json(payload: Any) -> str:
return json.dumps(payload, sort_keys=True, separators=(",", ":"), default=str)
def _sha256(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def plan_signature(plan: HydrocarbonStaplePlan) -> str:
"""Identity of a plan: everything that changes its terminal state."""
return _stable_json(
[
str(plan.block_id),
[int(plan.anchor_pair[0]), int(plan.anchor_pair[1])],
str(plan.ordered_pair),
int(plan.spacing),
[[int(position), str(monomer)] for position, monomer in plan.substitutions],
]
)
def plan_signatures(plans: Sequence[HydrocarbonStaplePlan]) -> list[str]:
return [plan_signature(plan) for plan in plans]
def _directory_digest(root: Path, suffixes: tuple[str, ...]) -> str:
"""Digest of (relative path, size, mtime) for weight files under ``root``.
Content hashing 400 MB of PeptiVerse weights on every run would cost more
than the cache saves, so identity is (path, size, mtime-ns). Touching or
swapping a weight file therefore invalidates the cache.
"""
if not root.is_dir():
return f"missing:{root}"
entries: list[tuple[str, int, int]] = []
for path in sorted(root.rglob("*")):
if not path.is_file() or path.suffix.lower() not in suffixes:
continue
stat = path.stat()
entries.append((str(path.relative_to(root)), int(stat.st_size), int(stat.st_mtime_ns)))
return _sha256(_stable_json(entries))
def _prior_digest(prior_dir: Path) -> str:
if not prior_dir.is_dir():
return f"missing:{prior_dir}"
entries = []
for path in sorted(prior_dir.rglob("*")):
if path.is_file() and path.suffix.lower() in (".json", ".csv", ".tsv", ".yaml", ".yml"):
stat = path.stat()
entries.append((str(path.relative_to(prior_dir)), int(stat.st_size), int(stat.st_mtime_ns)))
return _sha256(_stable_json(entries))
def build_fingerprint(
config: dict[str, Any],
*,
catalog: Iterable[Any],
repo_root: Path | None = None,
) -> dict[str, Any]:
"""Everything ``log q*`` depends on, other than the lead itself.
Any mismatch in any component invalidates cached entries automatically,
because the fingerprint hash is part of every row's key.
"""
repo_root = repo_root or Path.cwd()
hydro = dict(config.get("hydrocarbon") or {})
plan_control = dict(hydro.get("plan_control") or {})
terminal = dict(hydro.get("terminal_energy") or {})
property_cfg = dict(terminal.get("property") or {})
predictor = dict(config.get("property_predictor") or {})
priors = dict(config.get("reference_priors") or {})
peptide_prior = dict(priors.get("peptide") or {})
plan_reference = dict(hydro.get("plan_reference") or {})
mode_prior = dict(plan_reference.get("mode_prior") or {})
training = dict(config.get("training") or {})
# --- catalog identity -------------------------------------------------
catalog_entries = []
for block in catalog:
catalog_entries.append(
{
"block_id": getattr(block, "block_id", None),
"name": getattr(block, "name", None),
"chemistry_class": getattr(block, "chemistry_class", None),
"motif": getattr(block, "motif", None),
"ca_window": getattr(block, "ca_window", None),
"cost_score": getattr(block, "cost_score", None),
"spps_score": getattr(block, "spps_score", None),
}
)
prior_dir = mode_prior.get("prior_dir")
components = {
"schema_version": CACHE_SCHEMA_VERSION,
# --- exact-SB target definition ---
"exact_sb_beta": float(plan_control.get("exact_sb_beta", 1.0)),
"exact_sb_objective": bool(plan_control.get("exact_sb_objective", False)),
# Hard constraints define the support on which q_ref, q*, and q_theta
# are conditioned. Version this explicitly so pre-mask cache rows can
# never collide with post-mask targets.
"property_free_hard_plan_predicate": "v1",
"hard_plan_geometry_eps": float(
config.get("eps_geom", training.get("eps_geom", 2.5))
),
# --- catalog / empirical prior version ---
"catalog": catalog_entries,
"catalog_config": dict(hydro.get("catalog") or {}),
"mode_prior": mode_prior,
"mode_prior_digest": _prior_digest(
(repo_root / str(prior_dir)) if prior_dir else repo_root / "__missing__"
),
"plan_reference_enabled": bool(plan_reference.get("enabled", False)),
"plan_reference_bias": dict(plan_reference.get("bias") or {}),
"factorized_plan_reference": bool(
(hydro.get("reference") or {}).get("factorized_plan_reference", False)
),
# --- terminal energy / property config ---
"terminal_energy": {
key: terminal[key] for key in sorted(terminal) if key != "property"
},
"property": property_cfg,
"endpoint_prior": dict(hydro.get("endpoint_prior") or {}),
# These base-terminal coefficients feed E_T through base_terminal_factory.
"base_terminal_coefficients": {
key: training.get(key)
for key in ("lambda_close", "lambda_edit", "lambda_cost", "infeasible_penalty")
},
# --- PeptiVerse model / config ---
"peptiverse": {
key: predictor.get(key)
for key in (
"backend", "mode", "strict", "enable_fallback",
"allow_wt_token_fallback", "uncertainty", "offline",
"peptiverse_root", "classifier_weight_root", "manifest_path",
)
},
"peptiverse_manifest_digest": (
_sha256(Path(str(predictor["manifest_path"])).read_text())
if predictor.get("manifest_path") and Path(str(predictor["manifest_path"])).is_file()
else "missing_manifest"
),
"peptiverse_weights_digest": _directory_digest(
Path(str(predictor.get("classifier_weight_root", ""))) / "training_classifiers",
(".pt", ".json", ".joblib", ".bin", ".safetensors", ".txt"),
),
# --- reference priors (ESM-2 identity affects q_ref via the sampler) ---
"peptide_prior": {
key: peptide_prior.get(key)
for key in ("backend", "model_name_or_path", "temperature", "offline",
"strict_runtime", "ncaa_policy")
},
"anchor_prior": dict(priors.get("anchor") or {}),
"block_prior": dict(priors.get("block") or {}),
"reference_energy": dict(config.get("reference") or {}),
# --- geometry / edit settings ---
"geometry": dict(hydro.get("geometry") or {}),
"edit_constraints": dict(config.get("edit_constraints") or {}),
"curriculum_limits": {
key: (hydro.get("curriculum") or {}).get(key)
for key in ("max_anchor_edits", "prefer_existing_anchors", "protect_positions",
"require_valid_terminal")
},
"actions": dict(hydro.get("actions") or {}),
"max_neighbors": config.get("max_neighbors"),
"chemistry": config.get("chemistry"),
}
return {"hash": _sha256(_stable_json(components)), "components": components}
def lead_key(lead: Any) -> str:
"""Lead identity *and* content, so an edited sequence cannot reuse a row."""
return _stable_json(
{
"example_id": str(getattr(lead, "example_id", "")),
"linear_sequence": str(getattr(lead, "linear_sequence", "")),
"protected_positions": sorted(int(p) for p in (getattr(lead, "protected_positions", None) or [])),
# peptide_ca participates in the geometry term of E_T.
"peptide_ca": (getattr(lead, "target_context", None) or {}).get("peptide_ca"),
}
)
@dataclass
class ExactSBTargetEntry:
"""The deterministic half of the Exact-SB target for one lead."""
plan_signatures: list[str]
reference_logp: list[float]
terminal_energies: list[float]
log_q_star: list[float]
target_support_mask: list[bool] | None = None
joint_plan_count: int | None = None
def as_json(self) -> str:
payload: dict[str, Any] = {
"plan_signatures": self.plan_signatures,
"reference_logp": self.reference_logp,
"terminal_energies": self.terminal_energies,
"log_q_star": self.log_q_star,
}
# Preserve the byte-level shape of historical false-flag cache rows.
if self.target_support_mask is not None or self.joint_plan_count is not None:
payload["target_support_mask"] = self.target_support_mask
payload["joint_plan_count"] = self.joint_plan_count
return _stable_json(payload)
@classmethod
def from_json(cls, text: str) -> "ExactSBTargetEntry":
payload = json.loads(text)
return cls(
plan_signatures=list(payload["plan_signatures"]),
reference_logp=[float(v) for v in payload["reference_logp"]],
terminal_energies=[float(v) for v in payload["terminal_energies"]],
log_q_star=[float(v) for v in payload["log_q_star"]],
target_support_mask=(
[bool(v) for v in payload["target_support_mask"]]
if payload.get("target_support_mask") is not None
else None
),
joint_plan_count=(
int(payload["joint_plan_count"])
if payload.get("joint_plan_count") is not None
else None
),
)
@dataclass
class ExactSBCacheStats:
hits: int = 0
misses: int = 0
signature_mismatches: int = 0
writes: int = 0
plans_recomputed: int = 0
plans_served_from_cache: int = 0
energy_calls_saved: int = 0
def as_dict(self) -> dict[str, Any]:
total = self.hits + self.misses
return {
"hits": self.hits,
"misses": self.misses,
"lookups": total,
"hit_rate": (self.hits / total) if total else None,
"signature_mismatches": self.signature_mismatches,
"writes": self.writes,
"plans_recomputed": self.plans_recomputed,
"plans_served_from_cache": self.plans_served_from_cache,
"terminal_energy_calls_saved": self.energy_calls_saved,
}
class ExactSBTargetCache:
"""SQLite-backed store for per-lead Exact-SB targets.
SQLite (WAL, one row per lead) is used rather than one file per lead so a
4020-lead run does not create 4020 files, and so concurrent readers during
resume are safe. ``read_only=True`` gives a disabled/A-B arm that never
writes.
"""
def __init__(
self,
path: Path | str | None,
fingerprint: dict[str, Any],
*,
enabled: bool = True,
read_only: bool = False,
) -> None:
self.enabled = bool(enabled and path is not None)
self.read_only = bool(read_only)
self.fingerprint_hash = str(fingerprint["hash"])
self.fingerprint_components = fingerprint.get("components", {})
self.path = Path(path) if path is not None else None
self.stats = ExactSBCacheStats()
self._lock = threading.Lock()
self._connection: sqlite3.Connection | None = None
if self.enabled:
self._open()
# -- storage ---------------------------------------------------------
def _open(self) -> None:
assert self.path is not None
self.path.parent.mkdir(parents=True, exist_ok=True)
self._connection = sqlite3.connect(str(self.path), check_same_thread=False)
self._connection.execute("PRAGMA journal_mode=WAL")
self._connection.execute("PRAGMA synchronous=NORMAL")
self._connection.executescript(
"""
CREATE TABLE IF NOT EXISTS exact_sb_targets (
fingerprint TEXT NOT NULL,
lead_key TEXT NOT NULL,
payload TEXT NOT NULL,
n_plans INTEGER NOT NULL,
PRIMARY KEY (fingerprint, lead_key)
);
CREATE TABLE IF NOT EXISTS fingerprints (
fingerprint TEXT PRIMARY KEY,
components TEXT NOT NULL
);
"""
)
if not self.read_only:
self._connection.execute(
"INSERT OR REPLACE INTO fingerprints (fingerprint, components) VALUES (?, ?)",
(self.fingerprint_hash, _stable_json(self.fingerprint_components)),
)
self._connection.commit()
def close(self) -> None:
with self._lock:
if self._connection is not None:
self._connection.commit()
self._connection.close()
self._connection = None
# -- lookup / store --------------------------------------------------
def get(self, lead: Any, plans: Sequence[HydrocarbonStaplePlan]) -> ExactSBTargetEntry | None:
"""Return the cached target, or ``None`` on miss or plan drift."""
if not self.enabled or self._connection is None:
self.stats.misses += 1
return None
key = lead_key(lead)
with self._lock:
row = self._connection.execute(
"SELECT payload FROM exact_sb_targets WHERE fingerprint = ? AND lead_key = ?",
(self.fingerprint_hash, key),
).fetchone()
if row is None:
self.stats.misses += 1
return None
entry = ExactSBTargetEntry.from_json(row[0])
# Defence in depth: even on a fingerprint match, the live enumeration
# must produce exactly the same plans in the same order.
if entry.plan_signatures != plan_signatures(plans):
self.stats.signature_mismatches += 1
self.stats.misses += 1
return None
self.stats.hits += 1
self.stats.plans_served_from_cache += len(entry.plan_signatures)
self.stats.energy_calls_saved += len(entry.plan_signatures)
return entry
def put(self, lead: Any, entry: ExactSBTargetEntry) -> None:
if not self.enabled or self.read_only or self._connection is None:
return
with self._lock:
self._connection.execute(
"INSERT OR REPLACE INTO exact_sb_targets "
"(fingerprint, lead_key, payload, n_plans) VALUES (?, ?, ?, ?)",
(
self.fingerprint_hash,
lead_key(lead),
entry.as_json(),
len(entry.plan_signatures),
),
)
self._connection.commit()
self.stats.writes += 1
# -- diagnostics -----------------------------------------------------
def disk_bytes(self) -> int:
if self.path is None or not self.path.is_file():
return 0
total = self.path.stat().st_size
for suffix in ("-wal", "-shm"):
side = self.path.with_name(self.path.name + suffix)
if side.is_file():
total += side.stat().st_size
return int(total)
def row_count(self) -> int:
if not self.enabled or self._connection is None:
return 0
with self._lock:
return int(
self._connection.execute(
"SELECT COUNT(*) FROM exact_sb_targets WHERE fingerprint = ?",
(self.fingerprint_hash,),
).fetchone()[0]
)
def describe(self) -> dict[str, Any]:
return {
"enabled": self.enabled,
"read_only": self.read_only,
"path": None if self.path is None else str(self.path),
"fingerprint": self.fingerprint_hash,
"rows_for_fingerprint": self.row_count(),
"disk_bytes": self.disk_bytes(),
**self.stats.as_dict(),
}
def cache_path_from_config(config: dict[str, Any], repo_root: Path | None = None) -> Path | None:
"""Resolve ``hydrocarbon.plan_control.exact_sb_cache.path``; ``None`` = off."""
plan_control = dict(((config.get("hydrocarbon") or {}).get("plan_control") or {}))
section = dict(plan_control.get("exact_sb_cache") or {})
if not section or not bool(section.get("enabled", False)):
return None
raw = str(section.get("path") or "outputs/cache/exact_sb_targets.sqlite")
path = Path(raw)
if not path.is_absolute():
path = (repo_root or Path.cwd()) / path
return path
def energy_only_from_config(config: dict[str, Any]) -> bool:
"""Read ``hydrocarbon.plan_control.exact_sb_cache.energy_only`` (default on).
Off gives the original all-property scalar construction, which the A/B
benchmark uses as the reference arm.
"""
plan_control = dict(((config.get("hydrocarbon") or {}).get("plan_control") or {}))
section = dict(plan_control.get("exact_sb_cache") or {})
return bool(section.get("energy_only", True))
def build_cache_from_config(
config: dict[str, Any],
*,
catalog: Iterable[Any],
repo_root: Path | None = None,
read_only: bool = False,
override_path: Path | None = None,
) -> ExactSBTargetCache:
"""Construct the cache declared by ``config`` (disabled when absent)."""
catalog = list(catalog)
fingerprint = build_fingerprint(config, catalog=catalog, repo_root=repo_root)
path = override_path if override_path is not None else cache_path_from_config(config, repo_root)
return ExactSBTargetCache(path, fingerprint, enabled=path is not None, read_only=read_only)
def _accepts_energy_only(energy_fn: Any) -> bool:
"""Whether ``energy_fn`` takes the ``energy_only`` keyword."""
target = energy_fn.__call__ if not inspect.isfunction(energy_fn) else energy_fn
try:
signature = inspect.signature(target)
except (TypeError, ValueError):
return False
parameters = signature.parameters
if any(p.kind is inspect.Parameter.VAR_KEYWORD for p in parameters.values()):
return True
return "energy_only" in parameters
def resolve_exact_sb_target(
*,
lead: Any,
plans: Sequence[HydrocarbonStaplePlan],
reference_log_probabilities: torch.Tensor,
beta: float,
energy_fn: Any,
initial_state: Any,
build_terminal: Any,
cache: ExactSBTargetCache | None,
energy_only: bool = False,
scorer: Any = None,
scorer_config: Any = None,
) -> tuple[torch.Tensor, torch.Tensor, dict[str, Any]]:
"""Return ``(terminal_energies, log_q_star, info)`` for one lead.
On a cache hit the PeptiVerse-backed terminal energies are replayed from
disk. On a miss they are computed exactly as before and then stored.
``log q*`` is always produced by the shared
:func:`exact_sb_target_log_probabilities`, so the cached and recomputed
paths cannot diverge in definition.
``energy_only`` restricts property prediction to the properties that reach
the terminal energy and batch-prefetches their SMILES across all of this
lead's plans. It is a pure speedup: the energies, and therefore ``q*``, are
unchanged, so a cache row written by either path is interchangeable. It
requires ``scorer``/``scorer_config`` and an ``energy_fn`` accepting the
``energy_only`` keyword; without them the original scalar path runs.
"""
# Imported here to keep plan_control free of a dependency on this module.
from staplebridge.hydrocarbon.plan_control import (
exact_sb_target_log_probabilities,
)
from staplebridge.hydrocarbon.property_energy import required_energy_properties
device = reference_log_probabilities.device
dtype = reference_log_probabilities.dtype
if scorer_config is None:
scorer_config = getattr(energy_fn, "property_cfg", None)
joint_enabled = bool(
getattr(scorer_config, "enable_joint_perm_halflife_support", False)
)
def support_info(
mask_values: list[bool] | None, joint_count: int | None
) -> dict[str, Any]:
if not joint_enabled:
return {
"joint_perm_halflife_support_enabled": False,
"target_support_mask": None,
"joint_plan_count": None,
"joint_nonempty": None,
"joint_fallback": None,
"q_star_support_size": len(plans),
}
count = int(joint_count or 0)
fallback = count == 0
effective = None if fallback else list(mask_values or [])
return {
"joint_perm_halflife_support_enabled": True,
"target_support_mask": effective,
"joint_plan_count": count,
"joint_nonempty": not fallback,
"joint_fallback": fallback,
"q_star_support_size": count if count else len(plans),
}
entry = None if cache is None else cache.get(lead, plans)
if entry is not None:
if joint_enabled and entry.joint_plan_count is None:
# Backward-compatible payloads do not carry enough information to
# reconstruct a strict target support. A matching new-arm
# fingerprint should make this impossible, but fail closed.
entry = None
if entry is not None:
energies = torch.tensor(entry.terminal_energies, dtype=dtype, device=device)
mask = (
torch.tensor(entry.target_support_mask, dtype=torch.bool, device=device)
if joint_enabled
and int(entry.joint_plan_count or 0) > 0
and entry.target_support_mask is not None
else None
)
# Recomputed from the cached energies rather than trusting the stored
# log_q_star blindly; the stored copy is then verified against it.
log_q_star = exact_sb_target_log_probabilities(
reference_log_probabilities,
energies,
beta,
target_support_mask=mask,
)
stored = torch.tensor(entry.log_q_star, dtype=dtype, device=device)
finite = torch.isfinite(log_q_star) & torch.isfinite(stored)
max_drift = (
float((log_q_star[finite] - stored[finite]).abs().max().item())
if bool(finite.any().item())
else 0.0
)
if not torch.equal(torch.isneginf(log_q_star), torch.isneginf(stored)):
max_drift = float("inf")
return energies, log_q_star, {
"source": "cache",
"log_q_star_drift": max_drift,
**support_info(entry.target_support_mask, entry.joint_plan_count),
}
energies_list: list[float] = []
terminals = [build_terminal(initial_state, plan) for plan in plans]
# Energy-only + batched prefetch. Both are pure accelerations: the prefetch
# only warms caches, and energy_only skips predictions that are provably not
# summed into the energy. When either is unavailable the loop below is the
# original scalar path, so the stored energies are the same either way.
prefetch_info: dict[str, Any] = {}
if energy_only:
# Only the hydrocarbon terminal energy accepts the keyword and carries a
# scorer; any other callable (tests, lactam-style stubs) keeps the
# original scalar path rather than being handed an argument it rejects.
if not _accepts_energy_only(energy_fn):
energy_only = False
if energy_only:
# HydrocarbonTerminalEnergy carries both; taking them from energy_fn
# keeps the two knobs in sync with the energy that will actually run.
if scorer is None:
scorer = getattr(energy_fn, "property_scorer", None)
if scorer_config is None:
scorer_config = getattr(energy_fn, "property_cfg", None)
if scorer is None:
energy_only = False
if energy_only and scorer is not None:
properties = required_energy_properties(scorer_config) if scorer_config else ()
if properties:
smiles: list[str] = []
for terminal in terminals:
try:
smiles.extend(scorer.energy_only_smiles(terminal))
except Exception: # noqa: BLE001
# An unscorable terminal is handled by the energy itself
# (topology gate); nothing to prefetch for it.
continue
if smiles:
prefetch_info = scorer.prefetch(properties, smiles)
joint_mask_values: list[bool] = []
for terminal in terminals:
if energy_only:
energy, terms = energy_fn(initial_state, terminal, lead, energy_only=True)
else:
energy, terms = energy_fn(initial_state, terminal, lead)
energies_list.append(float(energy))
if joint_enabled:
if "hydrocarbon_joint_perm_halflife_condition" not in terms:
raise RuntimeError(
"joint Exact-SB support enabled but terminal energy did not "
"report the joint condition"
)
joint_mask_values.append(
bool(terms["hydrocarbon_joint_perm_halflife_condition"])
)
energies = torch.tensor(energies_list, dtype=dtype, device=device)
joint_count = sum(joint_mask_values) if joint_enabled else None
target_mask = (
torch.tensor(joint_mask_values, dtype=torch.bool, device=device)
if joint_enabled and int(joint_count or 0) > 0
else None
)
log_q_star = exact_sb_target_log_probabilities(
reference_log_probabilities,
energies,
beta,
target_support_mask=target_mask,
)
if cache is not None:
cache.stats.plans_recomputed += len(plans)
cache.put(
lead,
ExactSBTargetEntry(
plan_signatures=plan_signatures(plans),
reference_logp=[float(v) for v in reference_log_probabilities.detach().cpu().tolist()],
terminal_energies=energies_list,
log_q_star=[float(v) for v in log_q_star.detach().cpu().tolist()],
target_support_mask=(
list(joint_mask_values) if joint_enabled else None
),
joint_plan_count=joint_count,
),
)
return energies, log_q_star, {
"source": "computed",
"log_q_star_drift": 0.0,
"energy_only": bool(energy_only),
"prefetch": prefetch_info,
**support_info(
list(joint_mask_values) if joint_enabled else None,
joint_count,
),
}