| """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 |
|
|
| |
| |
| |
| DEFAULT_MODE_PRIOR_DIR: Final[str] = "staplebridge/hydrocarbon/data" |
|
|
| |
| |
| |
| |
| STAPLE_STRUCTURAL_EDIT_COST: Final[float] = 2.0 |
|
|
| |
| 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" |
|
|
| |
| 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.""" |
|
|
|
|
| |
| |
| |
|
|
|
|
| @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 |
| |
| |
| beta: float = 0.75 |
| |
| |
| 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", |
| } |
|
|
|
|
| |
| |
| |
|
|
|
|
| @dataclass |
| class PlanFilterConfig: |
| """Feasibility filters applied to enumerated plans.""" |
|
|
| |
| max_anchor_edits: int = 2 |
| |
| max_edit_budget: float = 6.0 |
| |
| 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: |
| |
| |
| 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] |
|
|
|
|
| |
| |
| |
|
|
|
|
| @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 |
| |
| |
| 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("-") |
|
|
| |
| 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 |
|
|
| |
| 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 |
| |
| |
| 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 |
| ) |
|
|
| |
| 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 |
|
|
| |
| 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: |
| |
| |
| 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] |
|
|
|
|
| |
| |
| |
|
|
|
|
| @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 |
| |
| 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 |
|
|
| |
| |
| |
| 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)) |
|
|