StapleBridge / staplebridge /hydrocarbon /plan_reference.py
pranamanam's picture Jingjie00's picture
Upload Staplebridge files (#1)
bb6d2aa
Raw
History Blame Contribute Delete
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
# ---------------------------------------------------------------------------
@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))