"""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, ), }