Download staplebridge/hydrocarbon/plan_reference.py from ChatterjeeLab/StapleBridge: direct link, hf CLI and curl.
- Browser
- Download file 36.9 kB
-
https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/hydrocarbon/plan_reference.py
- Command line
-
hf download hf://ChatterjeeLab/StapleBridge/staplebridge/hydrocarbon/plan_reference.py
-
curl -L -o plan_reference.py https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/hydrocarbon/plan_reference.py
36.9 kB
| """Plan-aware empirical hydrocarbon reference process. | |
| Why this module exists | |
| ---------------------- | |
| The hard-only reference in :mod:`staplebridge.hydrocarbon.actions` + | |
| :class:`staplebridge.reference.kernel.ReferenceKernel` reaches a stapled terminal | |
| on only ~21% of rollouts. The measured cause is **anchor overshoot**, not a | |
| scoring problem: the action generator offers an anchor-monomer substitution at | |
| almost every editable position, each individually legal, so an unguided walk | |
| installs 5-7 anchor monomers. ``validate_hydrocarbon_staple`` then returns | |
| ``DOUBLE_STAPLE_UNSUPPORTED``, anchor assignment is never offered, and the | |
| trajectory dead-ends with ``no_anchor_pair``. On a 32-lead probe, 98 of 101 | |
| failures had >2 anchors installed and no anchor pair. | |
| The fix is to commit to a *whole staple plan* before walking, then bias the walk | |
| toward finishing that plan: | |
| 1. enumerate every legal plan on the lead (S5-S5/i,i+4 and R8-S5/i,i+7); | |
| 2. filter on protected positions, anchor conflicts, edit budget and catalog; | |
| 3. draw one plan from q(plan | x) ∝ p_empirical(mode)^beta / n_mode(x); | |
| 4. bias the per-step kernel toward first anchor -> second anchor -> | |
| anchor/block assign -> topology activation for *that* plan; | |
| 5. downweight substitutions and anchor re-selection unrelated to the plan. | |
| The ``1 / n_mode(x)`` factor is the point of step 3: i,i+4 admits more anchor | |
| positions than i,i+7 on the same lead (8 vs 5 on a 12-mer), so weighting plans | |
| by the raw mode probability would amplify i,i+4 purely by opportunity count. | |
| Dividing by the per-lead legal-plan count of that mode makes the *mode* mass | |
| exactly ``p^beta`` and the choice *within* a mode uniform. | |
| The empirical mode prior is consumed **once, here, at plan selection**. It is | |
| deliberately not multiplied into every action and not re-counted in the terminal | |
| energy; ``configs/hydrocarbon_empirical_reference.yaml`` therefore sets | |
| ``endpoint_prior.weight_pair: 0.0`` so the same table cannot be charged twice. | |
| Isolation | |
| --------- | |
| Additive and hydrocarbon-only. Nothing here is imported by the lactam path: | |
| :class:`staplebridge.reference.kernel.ReferenceKernel`, | |
| :class:`staplebridge.reference.sampler.ReferenceTrajectorySampler`, | |
| ``staplebridge.graph.neighbors`` and ``BridgeTrainer`` are wrapped, never | |
| modified. The original hydrocarbon hard-only reference stays reachable exactly | |
| as before, so it remains available as the ablation baseline. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import random | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Any, Final | |
| import torch | |
| from staplebridge.chemistry.state import StapleState | |
| from staplebridge.data.schemas import BuildingBlock | |
| from staplebridge.hydrocarbon.catalog import block_topology, is_hydrocarbon_block | |
| from staplebridge.hydrocarbon.curriculum import ( | |
| HydrocarbonStaplePlan, | |
| propose_hydrocarbon_staple_plans, | |
| ) | |
| from staplebridge.hydrocarbon.factorized_plan_reference import ( | |
| FactorizedPlanReference, | |
| FactorizedPlanReferenceConfig, | |
| ) | |
| from staplebridge.hydrocarbon.tokenizer import is_anchor_token | |
| from staplebridge.reference.kernel import ReferenceKernel | |
| #: Versioned subset of the generated empirical priors needed for plan | |
| #: selection. Keeping it in the package makes defaults work in a clean clone; | |
| #: the complete analysis output remains optional and generated. | |
| DEFAULT_MODE_PRIOR_DIR: Final[str] = "staplebridge/hydrocarbon/data" | |
| #: Structural cost ``weighted_edit_distance`` charges for any completed staple: | |
| #: anchor 1.0 + topology 0.5 + block 0.5. A plan's terminal weighted edit | |
| #: distance is therefore ``n_edits + 2.0``, which is what the edit budget filter | |
| #: has to compare against. | |
| STAPLE_STRUCTURAL_EDIT_COST: Final[float] = 2.0 | |
| # -- action labels, relative to the committed plan --------------------------- | |
| ON_PLAN_FIRST_ANCHOR: Final[str] = "on_plan_first_anchor" | |
| ON_PLAN_SECOND_ANCHOR: Final[str] = "on_plan_second_anchor" | |
| ON_PLAN_ANCHOR_ASSIGN: Final[str] = "on_plan_anchor_assign" | |
| ON_PLAN_BLOCK_ASSIGN: Final[str] = "on_plan_block_assign" | |
| ON_PLAN_TOPOLOGY: Final[str] = "on_plan_topology_activation" | |
| OFF_PLAN_TOPOLOGY: Final[str] = "off_plan_topology_activation" | |
| OFF_PLAN_SUBSTITUTION: Final[str] = "off_plan_substitution" | |
| OFF_PLAN_ANCHOR: Final[str] = "off_plan_anchor_selection" | |
| OFF_PLAN_BLOCK: Final[str] = "off_plan_block_assign" | |
| PLAN_NOOP: Final[str] = "noop" | |
| #: Labels that count as progress on the committed plan. | |
| ON_PLAN_LABELS: Final[frozenset[str]] = frozenset( | |
| { | |
| ON_PLAN_FIRST_ANCHOR, | |
| ON_PLAN_SECOND_ANCHOR, | |
| ON_PLAN_ANCHOR_ASSIGN, | |
| ON_PLAN_BLOCK_ASSIGN, | |
| ON_PLAN_TOPOLOGY, | |
| } | |
| ) | |
| class PlanSelectionError(RuntimeError): | |
| """Raised when the empirical mode prior cannot be loaded.""" | |
| # --------------------------------------------------------------------------- | |
| # Empirical mode prior | |
| # --------------------------------------------------------------------------- | |
| class ModePriorConfig: | |
| """Config for :class:`EmpiricalModePrior`. | |
| Only the modes the catalog actually supports are kept, and their | |
| probabilities are renormalised over that restricted support. Without the | |
| renormalisation the ``beta`` exponent would act on a distribution whose mass | |
| partly sits on topologies the hard catalog forbids. | |
| """ | |
| prior_dir: str = DEFAULT_MODE_PRIOR_DIR | |
| dedup_version: str = "sequence_deduplicated" | |
| use_smoothed: bool = True | |
| #: Temperature on the empirical mode probabilities: ``p^beta``. 1.0 follows | |
| #: the data exactly, 0.0 is uniform over modes. | |
| beta: float = 0.75 | |
| #: Floor for a catalog mode absent from the table, so an enabled topology is | |
| #: never assigned probability zero. | |
| unobserved_probability: float = 1e-3 | |
| def from_dict(cls, data: dict[str, Any] | None) -> "ModePriorConfig": | |
| """Build from a ``hydrocarbon.plan_reference.mode_prior`` section.""" | |
| cfg = cls() | |
| for key, value in dict(data or {}).items(): | |
| if not hasattr(cfg, key): | |
| continue | |
| current = getattr(cfg, key) | |
| if isinstance(current, bool): | |
| setattr(cfg, key, bool(value)) | |
| elif isinstance(current, float): | |
| setattr(cfg, key, float(value)) | |
| else: | |
| setattr(cfg, key, value) | |
| return cfg | |
| class EmpiricalModePrior: | |
| """``p_empirical(mode)`` over the catalog's ``(pair, spacing)`` topologies. | |
| Args: | |
| catalog: the hydrocarbon blocks in play. Defines the support. | |
| config: prior configuration. | |
| root: repository root used to resolve a relative ``prior_dir``. | |
| Raises: | |
| PlanSelectionError: if the empirical table is missing or names no | |
| catalog mode. Failing loudly beats silently falling back to uniform, | |
| because "plan-aware *empirical* reference" would then be a misnomer. | |
| """ | |
| def __init__( | |
| self, | |
| catalog: list[BuildingBlock], | |
| config: ModePriorConfig | None = None, | |
| root: Path | None = None, | |
| ) -> None: | |
| self.cfg = config or ModePriorConfig() | |
| self._root = Path(root) if root is not None else Path(__file__).resolve().parents[2] | |
| self.modes: list[tuple[str, int]] = [ | |
| block_topology(b) for b in catalog if is_hydrocarbon_block(b) | |
| ] | |
| self._raw: dict[tuple[str, int], float] = {} | |
| self._probabilities: dict[tuple[str, int], float] = {} | |
| self._load() | |
| def prior_dir(self) -> Path: | |
| """Resolved directory holding the empirical JSON tables.""" | |
| candidate = Path(self.cfg.prior_dir) | |
| return candidate if candidate.is_absolute() else self._root / candidate | |
| def _load(self) -> None: | |
| """Read ``pair_spacing_probabilities.json`` and restrict to the catalog.""" | |
| path = self.prior_dir / "pair_spacing_probabilities.json" | |
| if not path.is_file(): | |
| raise PlanSelectionError( | |
| f"plan-aware reference needs the empirical mode table at {path}. " | |
| "It ships with this release at " | |
| "staplebridge/hydrocarbon/data/pair_spacing_probabilities.json; " | |
| "check hydrocarbon.plan_reference.mode_prior.prior_dir." | |
| ) | |
| with path.open("r", encoding="utf-8") as handle: | |
| payload = json.load(handle) | |
| versions = payload.get("probabilities_by_version") or {} | |
| if self.cfg.dedup_version not in versions: | |
| raise PlanSelectionError( | |
| f"dedup version {self.cfg.dedup_version!r} not in {path.name}; " | |
| f"available: {sorted(versions)}" | |
| ) | |
| categories = dict(versions[self.cfg.dedup_version].get("categories") or {}) | |
| field_name = ( | |
| "laplace_smoothed_probability" if self.cfg.use_smoothed else "raw_probability" | |
| ) | |
| for pair, spacing in self.modes: | |
| entry = categories.get(f"{pair}|{spacing}") or {} | |
| value = entry.get(field_name) | |
| self._raw[(pair, spacing)] = ( | |
| float(self.cfg.unobserved_probability) | |
| if value is None or float(value) <= 0.0 | |
| else float(value) | |
| ) | |
| total = sum(self._raw.values()) | |
| if total <= 0.0: | |
| raise PlanSelectionError( | |
| f"no catalog mode has positive empirical probability in {path.name}; " | |
| f"catalog modes: {self.modes}" | |
| ) | |
| self._probabilities = {k: v / total for k, v in self._raw.items()} | |
| def probability(self, mode: tuple[str, int]) -> float: | |
| """Renormalised ``p_empirical(mode)``; 0.0 for a non-catalog mode.""" | |
| return float(self._probabilities.get(mode, 0.0)) | |
| def tilted_weight(self, mode: tuple[str, int]) -> float: | |
| """``p_empirical(mode) ** beta``, the weight used at plan selection.""" | |
| probability = self.probability(mode) | |
| return 0.0 if probability <= 0.0 else probability ** float(self.cfg.beta) | |
| def describe(self) -> dict[str, Any]: | |
| """Summary for logging and audits.""" | |
| return { | |
| "prior_dir": str(self.prior_dir), | |
| "dedup_version": self.cfg.dedup_version, | |
| "use_smoothed": bool(self.cfg.use_smoothed), | |
| "beta": float(self.cfg.beta), | |
| "modes": [f"{p}/i,i+{s}" for p, s in self.modes], | |
| "p_empirical": { | |
| f"{p}/i,i+{s}": self.probability((p, s)) for p, s in self.modes | |
| }, | |
| "p_tilted": { | |
| f"{p}/i,i+{s}": self.tilted_weight((p, s)) for p, s in self.modes | |
| }, | |
| "uses_permeability_label": False, | |
| "is_trained_classifier": False, | |
| "consumed": "once, at plan selection", | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Plan enumeration, filtering and selection | |
| # --------------------------------------------------------------------------- | |
| class PlanFilterConfig: | |
| """Feasibility filters applied to enumerated plans.""" | |
| #: Reject plans needing more anchor substitutions than this. | |
| max_anchor_edits: int = 2 | |
| #: Terminal weighted-edit-distance ceiling (``edit_constraints.max_edit_budget``). | |
| max_edit_budget: float = 6.0 | |
| #: Minimum surviving sequence identity (``edit_constraints.min_sequence_identity``). | |
| min_sequence_identity: float = 0.60 | |
| def from_config( | |
| cls, hydro_cfg: dict[str, Any] | None, root_cfg: dict[str, Any] | None | |
| ) -> "PlanFilterConfig": | |
| """Read the curriculum and edit-constraint sections of a full config.""" | |
| curriculum = dict((hydro_cfg or {}).get("curriculum") or {}) | |
| edits = dict((root_cfg or {}).get("edit_constraints") or {}) | |
| return cls( | |
| max_anchor_edits=int(curriculum.get("max_anchor_edits", 2)), | |
| max_edit_budget=float(edits.get("max_edit_budget", 6.0)), | |
| min_sequence_identity=float(edits.get("min_sequence_identity", 0.60)), | |
| ) | |
| class PlanEnumerationReport: | |
| """Why plans were rejected, and what the surviving mode mix looks like. | |
| Every counter accumulates, so one report can be threaded through a whole | |
| batch of leads. ``n_enumerated`` and ``n_kept`` are therefore totals over all | |
| enumeration calls, not per-lead values — mixing the two conventions in one | |
| object would make the per-mode counts unreadable against them. | |
| """ | |
| n_calls: int = 0 | |
| n_enumerated: int = 0 | |
| n_kept: int = 0 | |
| rejected: dict[str, int] = field(default_factory=dict) | |
| per_mode_counts: dict[str, int] = field(default_factory=dict) | |
| def reject(self, reason: str) -> None: | |
| """Tally one rejection.""" | |
| self.rejected[reason] = self.rejected.get(reason, 0) + 1 | |
| def as_dict(self) -> dict[str, Any]: | |
| """JSON-serialisable view, with per-call means alongside the totals.""" | |
| calls = max(self.n_calls, 1) | |
| return { | |
| "n_calls": int(self.n_calls), | |
| "n_enumerated_total": int(self.n_enumerated), | |
| "n_kept_total": int(self.n_kept), | |
| "mean_enumerated_per_lead": float(self.n_enumerated / calls), | |
| "mean_kept_per_lead": float(self.n_kept / calls), | |
| "rejected": dict(sorted(self.rejected.items())), | |
| "per_mode_counts": dict(sorted(self.per_mode_counts.items())), | |
| } | |
| def enumerate_legal_plans( | |
| tokens: list[str], | |
| catalog: list[BuildingBlock], | |
| protected_positions: list[int] | None = None, | |
| filters: PlanFilterConfig | None = None, | |
| report: PlanEnumerationReport | None = None, | |
| ) -> list[HydrocarbonStaplePlan]: | |
| """Every legal staple plan on ``tokens``, after feasibility filtering. | |
| Delegates catalog/protected/anchor-conflict/double-staple filtering to | |
| :func:`~staplebridge.hydrocarbon.curriculum.propose_hydrocarbon_staple_plans` | |
| (so the plan-aware reference and the curriculum oracle agree on what is | |
| legal by construction), then applies the edit-budget and sequence-identity | |
| constraints the curriculum does not check. | |
| Returns: | |
| Plans in the curriculum's cheapest-first order. | |
| """ | |
| filters = filters or PlanFilterConfig() | |
| report = report if report is not None else PlanEnumerationReport() | |
| plans = propose_hydrocarbon_staple_plans( | |
| tokens, | |
| catalog, | |
| protected_positions=protected_positions, | |
| max_anchor_edits=filters.max_anchor_edits, | |
| ) | |
| report.n_calls += 1 | |
| report.n_enumerated += len(plans) | |
| kept: list[HydrocarbonStaplePlan] = [] | |
| for plan in plans: | |
| # Terminal weighted edit distance the plan would incur, including the | |
| # fixed structural cost of closing a staple. | |
| projected_edit = float(plan.n_edits) + STAPLE_STRUCTURAL_EDIT_COST | |
| if projected_edit > filters.max_edit_budget: | |
| report.reject("edit_budget_exhausted") | |
| continue | |
| identity = 1.0 - (plan.n_edits / len(tokens)) if tokens else 0.0 | |
| if identity < filters.min_sequence_identity: | |
| report.reject("below_min_sequence_identity") | |
| continue | |
| kept.append(plan) | |
| mode = f"{plan.ordered_pair}/i,i+{plan.spacing}" | |
| report.per_mode_counts[mode] = report.per_mode_counts.get(mode, 0) + 1 | |
| report.n_kept += len(kept) | |
| return kept | |
| def plan_selection_weights( | |
| plans: list[HydrocarbonStaplePlan], mode_prior: EmpiricalModePrior | |
| ) -> list[float]: | |
| """``q(plan | x) ∝ p_empirical(mode)^beta / n_mode(x)``, unnormalised. | |
| Dividing by ``n_mode(x)`` — the number of legal plans of that mode *on this | |
| lead* — is what keeps i,i+4 from being amplified simply because it has more | |
| admissible anchor positions than i,i+7. The resulting mode marginal is | |
| exactly ``p^beta`` and the within-mode choice is uniform. | |
| """ | |
| counts: dict[tuple[str, int], int] = {} | |
| for plan in plans: | |
| key = (plan.ordered_pair, plan.spacing) | |
| counts[key] = counts.get(key, 0) + 1 | |
| weights: list[float] = [] | |
| for plan in plans: | |
| key = (plan.ordered_pair, plan.spacing) | |
| n_mode = counts[key] | |
| weights.append(mode_prior.tilted_weight(key) / float(n_mode) if n_mode else 0.0) | |
| return weights | |
| def select_plan( | |
| plans: list[HydrocarbonStaplePlan], | |
| mode_prior: EmpiricalModePrior, | |
| rng: random.Random, | |
| ) -> HydrocarbonStaplePlan | None: | |
| """Draw one plan from ``q(plan | x)``. | |
| Returns ``None`` when there is no legal plan, or when every legal plan's mode | |
| has zero empirical weight. | |
| """ | |
| if not plans: | |
| return None | |
| weights = plan_selection_weights(plans, mode_prior) | |
| total = sum(weights) | |
| if total <= 0.0: | |
| return None | |
| threshold = rng.random() * total | |
| cumulative = 0.0 | |
| for plan, weight in zip(plans, weights): | |
| cumulative += weight | |
| if cumulative >= threshold: | |
| return plan | |
| return plans[-1] | |
| # --------------------------------------------------------------------------- | |
| # Plan-conditional action labelling and biasing | |
| # --------------------------------------------------------------------------- | |
| class PlanBiasConfig: | |
| """Log-space bonuses applied to the reference pmf, per plan-relative label. | |
| Positive values favour an action, negative values suppress it. The four | |
| on-plan structural bonuses increase along the build order (first anchor -> | |
| second anchor -> assign -> activate) so that a partially built plan is | |
| always pulled forward rather than left to compete with a fresh restart. | |
| The off-plan substitution penalty is the load-bearing one: the action | |
| generator offers an anchor substitution at nearly every editable position, | |
| and unguided that is what installs a third anchor and kills the trajectory. | |
| """ | |
| first_anchor: float = 3.0 | |
| second_anchor: float = 3.5 | |
| anchor_assign: float = 4.0 | |
| block_assign: float = 4.0 | |
| topology_activation: float = 4.5 | |
| #: Closing a different pair contradicts the committed plan and strict | |
| #: hierarchical inference. Keep it a failure, not an alternative positive. | |
| off_plan_topology_activation: float = -4.5 | |
| off_plan_substitution: float = -3.0 | |
| off_plan_anchor_selection: float = -3.0 | |
| off_plan_block_assign: float = -1.0 | |
| noop: float = -1.0 | |
| def from_dict(cls, data: dict[str, Any] | None) -> "PlanBiasConfig": | |
| """Build from a ``hydrocarbon.plan_reference.bias`` section.""" | |
| cfg = cls() | |
| for key, value in dict(data or {}).items(): | |
| if hasattr(cfg, key): | |
| setattr(cfg, key, float(value)) | |
| return cfg | |
| def as_dict(self) -> dict[str, float]: | |
| """Label -> bonus mapping used by the kernel.""" | |
| return { | |
| ON_PLAN_FIRST_ANCHOR: self.first_anchor, | |
| ON_PLAN_SECOND_ANCHOR: self.second_anchor, | |
| ON_PLAN_ANCHOR_ASSIGN: self.anchor_assign, | |
| ON_PLAN_BLOCK_ASSIGN: self.block_assign, | |
| ON_PLAN_TOPOLOGY: self.topology_activation, | |
| OFF_PLAN_TOPOLOGY: self.off_plan_topology_activation, | |
| OFF_PLAN_SUBSTITUTION: self.off_plan_substitution, | |
| OFF_PLAN_ANCHOR: self.off_plan_anchor_selection, | |
| OFF_PLAN_BLOCK: self.off_plan_block_assign, | |
| PLAN_NOOP: self.noop, | |
| } | |
| def plan_positions_satisfied( | |
| tokens: list[str], plan: HydrocarbonStaplePlan | |
| ) -> tuple[bool, bool]: | |
| """Whether the plan's ``i`` and ``j`` anchor monomers are already installed.""" | |
| i, j = plan.anchor_pair | |
| i_token, j_token = plan.ordered_pair.split("-") | |
| have_i = 0 <= i < len(tokens) and tokens[i].upper() == i_token | |
| have_j = 0 <= j < len(tokens) and tokens[j].upper() == j_token | |
| return have_i, have_j | |
| def classify_against_plan( | |
| state: StapleState, candidate: StapleState, plan: HydrocarbonStaplePlan | |
| ) -> str: | |
| """Label the transition ``state -> candidate`` relative to ``plan``. | |
| Checked in the same order the build proceeds, so a composite transition | |
| (the action generator sets ``block_id`` in the same step as the anchor | |
| assignment) is attributed to its most advanced effect. | |
| """ | |
| plan_i, plan_j = plan.anchor_pair | |
| i_token, j_token = plan.ordered_pair.split("-") | |
| # -- topology activation -------------------------------------------- | |
| if state.topology != candidate.topology: | |
| if candidate.topology != "stapled": | |
| return PLAN_NOOP | |
| on_plan = ( | |
| candidate.anchor_pair is not None | |
| and tuple(candidate.anchor_pair) == (plan_i, plan_j) | |
| and candidate.block_id == plan.block_id | |
| ) | |
| return ON_PLAN_TOPOLOGY if on_plan else OFF_PLAN_TOPOLOGY | |
| # -- sequence edit --------------------------------------------------- | |
| if state.sequence_tokens != candidate.sequence_tokens: | |
| changed = [ | |
| position | |
| for position in range(min(len(state.sequence_tokens), len(candidate.sequence_tokens))) | |
| if state.sequence_tokens[position] != candidate.sequence_tokens[position] | |
| ] | |
| if len(changed) != 1: | |
| return OFF_PLAN_SUBSTITUTION | |
| position = changed[0] | |
| installed = candidate.sequence_tokens[position].upper() | |
| wanted = ( | |
| i_token if position == plan_i else j_token if position == plan_j else None | |
| ) | |
| if wanted is None or installed != wanted: | |
| return OFF_PLAN_SUBSTITUTION | |
| # Ordering is by *progress*, not by index: whichever of the two plan | |
| # anchors lands first is the "first anchor" install. | |
| have_i, have_j = plan_positions_satisfied(state.sequence_tokens, plan) | |
| return ( | |
| ON_PLAN_SECOND_ANCHOR if (have_i or have_j) else ON_PLAN_FIRST_ANCHOR | |
| ) | |
| # -- anchor selection ------------------------------------------------ | |
| if state.anchor_pair != candidate.anchor_pair: | |
| if ( | |
| candidate.anchor_pair is not None | |
| and tuple(candidate.anchor_pair) == (plan_i, plan_j) | |
| and candidate.block_id in (None, plan.block_id) | |
| ): | |
| return ON_PLAN_ANCHOR_ASSIGN | |
| return OFF_PLAN_ANCHOR | |
| # -- block assignment ------------------------------------------------ | |
| if state.block_id != candidate.block_id: | |
| if ( | |
| candidate.block_id == plan.block_id | |
| and candidate.anchor_pair is not None | |
| and tuple(candidate.anchor_pair) == (plan_i, plan_j) | |
| ): | |
| return ON_PLAN_BLOCK_ASSIGN | |
| return OFF_PLAN_BLOCK | |
| return PLAN_NOOP | |
| class PlanAwareReferenceKernel: | |
| """Reference kernel that conditions on a committed staple plan. | |
| Wraps an unmodified :class:`~staplebridge.reference.kernel.ReferenceKernel`: | |
| the base pmf (peptide/anchor/block priors, cost, geometry, action-progress, | |
| group normalisation, substitution downweight) is computed exactly as today, | |
| then reweighted by ``exp(bonus(label))`` and renormalised. With no plan | |
| committed, or with all bonuses at zero, this is the base kernel. | |
| Reweighting in probability space rather than editing the base logits keeps | |
| the two kernels directly comparable for the ablation: the only difference is | |
| a plan-conditional multiplicative factor. | |
| """ | |
| def __init__( | |
| self, base_kernel: ReferenceKernel, bias: PlanBiasConfig | None = None | |
| ) -> None: | |
| self.base_kernel = base_kernel | |
| self.bias = bias or PlanBiasConfig() | |
| self._bonuses = self.bias.as_dict() | |
| def labels( | |
| self, | |
| state: StapleState, | |
| candidates: list[StapleState], | |
| plan: HydrocarbonStaplePlan | None, | |
| ) -> list[str]: | |
| """Plan-relative label for each candidate.""" | |
| if plan is None: | |
| return [PLAN_NOOP] * len(candidates) | |
| return [classify_against_plan(state, c, plan) for c in candidates] | |
| def plan_probs( | |
| self, | |
| state: StapleState, | |
| candidates: list[StapleState], | |
| plan: HydrocarbonStaplePlan | None, | |
| context: dict[str, Any] | None = None, | |
| ) -> tuple[torch.Tensor, list[str]]: | |
| """Plan-conditional pmf over ``candidates``, plus their labels.""" | |
| probs = self.base_kernel.reference_probs(state, candidates, context=context) | |
| if plan is None: | |
| return probs, [PLAN_NOOP] * len(candidates) | |
| labels = self.labels(state, candidates, plan) | |
| factors = torch.tensor( | |
| [math.exp(self._bonuses.get(label, 0.0)) for label in labels], | |
| dtype=torch.float32, | |
| ) | |
| tilted = probs * factors | |
| total = float(tilted.sum().item()) | |
| if total <= 0.0: | |
| # Every candidate had zero base mass; fall back rather than emit a | |
| # degenerate pmf that ``torch.multinomial`` would reject. | |
| return probs, labels | |
| return tilted / total, labels | |
| def sample_next( | |
| self, | |
| state: StapleState, | |
| candidates: list[StapleState], | |
| plan: HydrocarbonStaplePlan | None, | |
| context: dict[str, Any] | None = None, | |
| ) -> tuple[StapleState, str]: | |
| """Draw one candidate from the plan-conditional pmf.""" | |
| probs, labels = self.plan_probs(state, candidates, plan, context=context) | |
| index = int(torch.multinomial(probs, num_samples=1).item()) | |
| return candidates[index], labels[index] | |
| # --------------------------------------------------------------------------- | |
| # Plan-aware trajectory sampler | |
| # --------------------------------------------------------------------------- | |
| class PlanProgress: | |
| """Which stages of the committed plan a trajectory actually reached. | |
| Recorded per stage rather than as a single success flag, because when the | |
| stapled rate disappoints the question is always *which* stage lost the | |
| trajectory. | |
| """ | |
| plan_selected: bool = False | |
| first_anchor_installed: bool = False | |
| second_anchor_installed: bool = False | |
| anchor_assigned: bool = False | |
| block_assigned: bool = False | |
| topology_activated: bool = False | |
| plan_completed: bool = False | |
| n_on_plan_actions: int = 0 | |
| n_off_plan_substitutions: int = 0 | |
| n_off_plan_anchor_selections: int = 0 | |
| n_actions: int = 0 | |
| def unrelated_substitution_rate(self) -> float: | |
| """Share of this trajectory's actions that were off-plan substitutions.""" | |
| return ( | |
| self.n_off_plan_substitutions / self.n_actions if self.n_actions else 0.0 | |
| ) | |
| def as_dict(self) -> dict[str, Any]: | |
| """JSON-serialisable view.""" | |
| return { | |
| "plan_selected": bool(self.plan_selected), | |
| "first_anchor_installed": bool(self.first_anchor_installed), | |
| "second_anchor_installed": bool(self.second_anchor_installed), | |
| "anchor_assigned": bool(self.anchor_assigned), | |
| "block_assigned": bool(self.block_assigned), | |
| "topology_activated": bool(self.topology_activated), | |
| "plan_completed": bool(self.plan_completed), | |
| "n_on_plan_actions": int(self.n_on_plan_actions), | |
| "n_off_plan_substitutions": int(self.n_off_plan_substitutions), | |
| "n_off_plan_anchor_selections": int(self.n_off_plan_anchor_selections), | |
| "n_actions": int(self.n_actions), | |
| "unrelated_substitution_rate": float(self.unrelated_substitution_rate), | |
| } | |
| class PlanAwareTrajectory: | |
| """One plan-aware rollout.""" | |
| states: list[StapleState] | |
| plan: HydrocarbonStaplePlan | None | |
| progress: PlanProgress | |
| action_labels: list[str] = field(default_factory=list) | |
| no_plan_reason: str | None = None | |
| #: True when the rollout stopped because the graph offered no neighbour. | |
| no_neighbor: bool = False | |
| class PlanAwareReferenceSampler: | |
| """Reference sampler that commits to a plan, then completes it. | |
| Args: | |
| graph: the hydrocarbon transition graph (unmodified). | |
| kernel: the plan-conditional kernel. | |
| mode_prior: empirical mode prior, consumed once per trajectory. | |
| filters: plan feasibility filters. | |
| seed: base seed for plan selection, kept separate from the global torch | |
| RNG so plan draws are reproducible independently of the pmf draws. | |
| """ | |
| def __init__( | |
| self, | |
| graph: Any, | |
| kernel: PlanAwareReferenceKernel, | |
| mode_prior: EmpiricalModePrior, | |
| filters: PlanFilterConfig | None = None, | |
| seed: int = 42, | |
| factorized_reference: FactorizedPlanReference | None = None, | |
| ) -> None: | |
| self.graph = graph | |
| self.kernel = kernel | |
| self.mode_prior = mode_prior | |
| self.filters = filters or PlanFilterConfig() | |
| self.factorized_reference = factorized_reference | |
| self._rng = random.Random(seed) | |
| def factorized_plan_reference_enabled(self) -> bool: | |
| return self.factorized_reference is not None | |
| def plan_selection_weights( | |
| self, | |
| initial: StapleState, | |
| plans: list[HydrocarbonStaplePlan], | |
| context: dict[str, Any] | None = None, | |
| ) -> list[float]: | |
| """Active plan-reference weights, with a bit-exact legacy branch.""" | |
| if self.factorized_reference is None: | |
| return plan_selection_weights(plans, self.mode_prior) | |
| return self.factorized_reference.weights(initial, plans, context) | |
| def select_plan( | |
| self, | |
| initial: StapleState, | |
| plans: list[HydrocarbonStaplePlan], | |
| context: dict[str, Any] | None = None, | |
| ) -> HydrocarbonStaplePlan | None: | |
| if not plans: | |
| return None | |
| weights = self.plan_selection_weights(initial, plans, context) | |
| total = sum(weights) | |
| if total <= 0.0: | |
| return None | |
| threshold = self._rng.random() * total | |
| cumulative = 0.0 | |
| for plan, weight in zip(plans, weights): | |
| cumulative += weight | |
| if cumulative >= threshold: | |
| return plan | |
| return plans[-1] | |
| def sample_trajectory( | |
| self, | |
| init_state: StapleState, | |
| protected_positions: list[int], | |
| context: dict[str, Any], | |
| horizon: int, | |
| early_stop: bool = True, | |
| report: PlanEnumerationReport | None = None, | |
| ) -> PlanAwareTrajectory: | |
| """Select a plan for ``init_state``, then walk toward completing it.""" | |
| plans = enumerate_legal_plans( | |
| init_state.sequence_tokens, | |
| self.graph.catalog, | |
| protected_positions=protected_positions, | |
| filters=self.filters, | |
| report=report, | |
| ) | |
| plan = self.select_plan(init_state, plans, context) | |
| progress = PlanProgress(plan_selected=plan is not None) | |
| if plan is None: | |
| reason = "no_legal_plan" if not plans else "no_mode_weight" | |
| return PlanAwareTrajectory( | |
| states=[init_state], plan=None, progress=progress, no_plan_reason=reason | |
| ) | |
| states = [init_state] | |
| labels: list[str] = [] | |
| current = init_state | |
| no_neighbor = False | |
| for _ in range(horizon): | |
| candidates = self.graph.neighbors( | |
| current, protected_positions=protected_positions | |
| ) | |
| if not candidates: | |
| no_neighbor = True | |
| break | |
| nxt, label = self.kernel.sample_next( | |
| current, candidates, plan, context=context | |
| ) | |
| states.append(nxt) | |
| labels.append(label) | |
| progress.n_actions += 1 | |
| if label in ON_PLAN_LABELS: | |
| progress.n_on_plan_actions += 1 | |
| if label == ON_PLAN_FIRST_ANCHOR: | |
| progress.first_anchor_installed = True | |
| elif label == ON_PLAN_SECOND_ANCHOR: | |
| progress.second_anchor_installed = True | |
| elif label == ON_PLAN_ANCHOR_ASSIGN: | |
| progress.anchor_assigned = True | |
| elif label == ON_PLAN_BLOCK_ASSIGN: | |
| progress.block_assigned = True | |
| elif label == ON_PLAN_TOPOLOGY: | |
| progress.topology_activated = True | |
| elif label == OFF_PLAN_TOPOLOGY: | |
| progress.topology_activated = True | |
| elif label == OFF_PLAN_SUBSTITUTION: | |
| progress.n_off_plan_substitutions += 1 | |
| elif label == OFF_PLAN_ANCHOR: | |
| progress.n_off_plan_anchor_selections += 1 | |
| current = nxt | |
| if early_stop and current.topology == "stapled": | |
| break | |
| # The anchor assignment is composite (it sets block_id in the same | |
| # transition), so credit block assignment from the terminal state rather | |
| # than requiring a separate labelled step. | |
| if current.block_id == plan.block_id and tuple( | |
| current.anchor_pair or (-1, -1) | |
| ) == plan.anchor_pair: | |
| progress.block_assigned = True | |
| have_i, have_j = plan_positions_satisfied(current.sequence_tokens, plan) | |
| if have_i and have_j: | |
| progress.first_anchor_installed = True | |
| progress.second_anchor_installed = True | |
| elif have_i or have_j: | |
| progress.first_anchor_installed = True | |
| progress.plan_completed = bool( | |
| current.topology == "stapled" | |
| and current.anchor_pair is not None | |
| and tuple(current.anchor_pair) == plan.anchor_pair | |
| and current.block_id == plan.block_id | |
| ) | |
| return PlanAwareTrajectory( | |
| states=states, | |
| plan=plan, | |
| progress=progress, | |
| action_labels=labels, | |
| no_neighbor=no_neighbor, | |
| ) | |
| def sample_batch( | |
| self, | |
| init_state: StapleState, | |
| protected_positions: list[int], | |
| context: dict[str, Any], | |
| horizon: int, | |
| n: int, | |
| report: PlanEnumerationReport | None = None, | |
| ) -> list[PlanAwareTrajectory]: | |
| """``n`` independent plan-aware rollouts from ``init_state``.""" | |
| return [ | |
| self.sample_trajectory( | |
| init_state, | |
| protected_positions=protected_positions, | |
| context=context, | |
| horizon=horizon, | |
| report=report, | |
| ) | |
| for _ in range(n) | |
| ] | |
| class PlanReferenceConfig: | |
| """Full config for the plan-aware reference, from a ``hydrocarbon`` section.""" | |
| enabled: bool = False | |
| mode_prior: ModePriorConfig = field(default_factory=ModePriorConfig) | |
| bias: PlanBiasConfig = field(default_factory=PlanBiasConfig) | |
| filters: PlanFilterConfig = field(default_factory=PlanFilterConfig) | |
| factorized: FactorizedPlanReferenceConfig = field( | |
| default_factory=FactorizedPlanReferenceConfig | |
| ) | |
| def from_config(cls, root_cfg: dict[str, Any] | None) -> "PlanReferenceConfig": | |
| """Read ``hydrocarbon.plan_reference`` plus the shared edit constraints.""" | |
| root_cfg = dict(root_cfg or {}) | |
| hydro_cfg = dict(root_cfg.get("hydrocarbon") or {}) | |
| section = dict(hydro_cfg.get("plan_reference") or {}) | |
| return cls( | |
| enabled=bool(section.get("enabled", False)), | |
| mode_prior=ModePriorConfig.from_dict(section.get("mode_prior")), | |
| bias=PlanBiasConfig.from_dict(section.get("bias")), | |
| filters=PlanFilterConfig.from_config(hydro_cfg, root_cfg), | |
| factorized=FactorizedPlanReferenceConfig.from_config(root_cfg), | |
| ) | |
| def build_plan_aware_sampler( | |
| graph: Any, | |
| base_kernel: ReferenceKernel, | |
| root_cfg: dict[str, Any] | None, | |
| seed: int = 42, | |
| root: Path | None = None, | |
| ) -> tuple[PlanAwareReferenceSampler, PlanReferenceConfig]: | |
| """Assemble the plan-aware sampler from a full config mapping.""" | |
| cfg = PlanReferenceConfig.from_config(root_cfg) | |
| mode_prior = EmpiricalModePrior(graph.catalog, cfg.mode_prior, root=root) | |
| kernel = PlanAwareReferenceKernel(base_kernel, cfg.bias) | |
| factorized_reference = None | |
| if cfg.factorized.enabled: | |
| energy = base_kernel.energy_model | |
| factorized_reference = FactorizedPlanReference( | |
| mode_prior=mode_prior, | |
| catalog=graph.catalog, | |
| peptide_prior=energy.peptide_prior, | |
| anchor_prior=energy.anchor_prior, | |
| block_prior=energy.block_prior, | |
| config=cfg.factorized, | |
| ) | |
| sampler = PlanAwareReferenceSampler( | |
| graph, | |
| kernel, | |
| mode_prior, | |
| filters=cfg.filters, | |
| seed=seed, | |
| factorized_reference=factorized_reference, | |
| ) | |
| return sampler, cfg | |
| def count_anchor_monomers(tokens: list[str]) -> int: | |
| """Number of hydrocarbon anchor monomers in ``tokens``.""" | |
| return sum(1 for t in tokens if is_anchor_token(t)) | |