Download staplebridge/hydrocarbon/exact_sb_cache.py from ChatterjeeLab/StapleBridge: direct link, hf CLI and curl.
- Browser
- Download file 28.6 kB
-
https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/hydrocarbon/exact_sb_cache.py
- Command line
-
hf download hf://ChatterjeeLab/StapleBridge/staplebridge/hydrocarbon/exact_sb_cache.py
-
curl -L -o exact_sb_cache.py https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/hydrocarbon/exact_sb_cache.py
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"), | |
| } | |
| ) | |
| 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) | |
| 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 | |
| ), | |
| ) | |
| 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, | |
| ), | |
| } | |