File size: 36,924 Bytes
bb6d2aa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 | """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))
|