"""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 # --------------------------------------------------------------------------- @dataclass 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 @classmethod 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() @property 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 # --------------------------------------------------------------------------- @dataclass 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 @classmethod 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)), ) @dataclass 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 # --------------------------------------------------------------------------- @dataclass 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 @classmethod 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 # --------------------------------------------------------------------------- @dataclass 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 @property 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), } @dataclass 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) @property 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) ] @dataclass 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 ) @classmethod 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))