| """Topology-mass-preserving hydrocarbon plan reference. |
| |
| This module is hydrocarbon-only. It does not import or modify the lactam |
| catalog, decoder, property model, SMILES builder, plan space, or loss. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import math |
| from dataclasses import dataclass |
| from typing import Any |
|
|
| from staplebridge.chemistry.state import StapleState |
| from staplebridge.data.schemas import BuildingBlock |
| from staplebridge.hydrocarbon.curriculum import ( |
| HydrocarbonStaplePlan, |
| build_hydrocarbon_demonstration_path, |
| ) |
|
|
|
|
| @dataclass |
| class FactorizedPlanReferenceConfig: |
| """Configuration read from ``hydrocarbon.reference``. |
| |
| Defaults were preregistered without property labels from the component |
| scale audit: MotifSupportAnchorPrior/geometry is primary and frozen ESM2 |
| delta is only a weak regularizer. CatalogBlockPrior is a legality check and |
| diagnostic, not a soft ranking term. |
| """ |
|
|
| enabled: bool = False |
| geometry_coefficient: float = 1.0 |
| esm2_coefficient: float = 0.1 |
| temperature: float = 1.0 |
| block_legality_only: bool = True |
|
|
| @classmethod |
| def from_config( |
| cls, root_cfg: dict[str, Any] | None |
| ) -> "FactorizedPlanReferenceConfig": |
| root_cfg = dict(root_cfg or {}) |
| hydro = dict(root_cfg.get("hydrocarbon") or {}) |
| section = dict(hydro.get("reference") or {}) |
| within = dict(section.get("within_mode") or {}) |
| return cls( |
| enabled=bool(section.get("factorized_plan_reference", False)), |
| geometry_coefficient=float(within.get("geometry_coefficient", 1.0)), |
| esm2_coefficient=float(within.get("esm2_coefficient", 0.1)), |
| temperature=max(float(within.get("temperature", 1.0)), 1e-8), |
| block_legality_only=bool(within.get("block_legality_only", True)), |
| ) |
|
|
| def describe(self) -> dict[str, Any]: |
| return { |
| "factorized_plan_reference": bool(self.enabled), |
| "geometry_coefficient": float(self.geometry_coefficient), |
| "esm2_coefficient": float(self.esm2_coefficient), |
| "temperature": float(self.temperature), |
| "block_legality_only": bool(self.block_legality_only), |
| "uses_property_labels": False, |
| "normalization": "softmax separately within each feasible mode", |
| } |
|
|
|
|
| class FactorizedPlanReference: |
| """Compute ``q_mode(mode|lead) * q_within(plan|lead,mode)``. |
| |
| The supplied ``mode_prior`` owns the StaPep probability and empirical |
| beta. Within-mode components are normalized separately, so they cannot |
| change the total mass assigned to a topology. |
| """ |
|
|
| def __init__( |
| self, |
| mode_prior: Any, |
| catalog: list[BuildingBlock], |
| peptide_prior: Any, |
| anchor_prior: Any, |
| block_prior: Any, |
| config: FactorizedPlanReferenceConfig, |
| ) -> None: |
| self.mode_prior = mode_prior |
| self.catalog = list(catalog) |
| self.catalog_index = {block.block_id: block for block in catalog} |
| self.peptide_prior = peptide_prior |
| self.anchor_prior = anchor_prior |
| self.block_prior = block_prior |
| self.cfg = config |
| self.last_diagnostics: list[dict[str, Any]] = [] |
|
|
| @staticmethod |
| def _softmax(values: list[float]) -> list[float]: |
| if not values: |
| return [] |
| peak = max(values) |
| exponentials = [math.exp(value - peak) for value in values] |
| total = sum(exponentials) |
| if total <= 0.0 or not math.isfinite(total): |
| return [1.0 / len(values)] * len(values) |
| probabilities = [value / total for value in exponentials] |
| if len(probabilities) > 1: |
| probabilities[-1] = 1.0 - sum(probabilities[:-1]) |
| return probabilities |
|
|
| def weights( |
| self, |
| initial: StapleState, |
| plans: list[HydrocarbonStaplePlan], |
| context: dict[str, Any] | None = None, |
| ) -> list[float]: |
| """Return normalized factorized probabilities in ``plans`` order.""" |
| if not plans: |
| self.last_diagnostics = [] |
| return [] |
|
|
| context = dict(context or {}) |
| terminals = [ |
| build_hydrocarbon_demonstration_path(initial, plan, self.catalog)[-1] |
| for plan in plans |
| ] |
| esm2_scores = self.peptide_prior.batch_score_transitions( |
| initial, terminals, context |
| ) |
|
|
| feasible_modes: list[tuple[str, int]] = [] |
| for plan in plans: |
| mode = (plan.ordered_pair, plan.spacing) |
| if mode not in feasible_modes: |
| feasible_modes.append(mode) |
| tilted = [self.mode_prior.tilted_weight(mode) for mode in feasible_modes] |
| tilted_total = sum(tilted) |
| if tilted_total <= 0.0: |
| self.last_diagnostics = [] |
| return [0.0] * len(plans) |
| q_mode = { |
| mode: weight / tilted_total for mode, weight in zip(feasible_modes, tilted) |
| } |
| if len(feasible_modes) > 1: |
| q_mode[feasible_modes[-1]] = 1.0 - sum( |
| q_mode[mode] for mode in feasible_modes[:-1] |
| ) |
|
|
| logits: list[float] = [] |
| diagnostics: list[dict[str, Any]] = [] |
| for plan, terminal, esm2_score in zip(plans, terminals, esm2_scores): |
| block = self.catalog_index.get(plan.block_id) |
| if block is None: |
| anchor_score = float("-inf") |
| block_score = float("-inf") |
| anchor_components: dict[str, float] = {} |
| legal = False |
| else: |
| anchor_context = dict(context) |
| anchor_context["return_components"] = True |
| anchor_score = float( |
| self.anchor_prior.score_anchor( |
| terminal.sequence_tokens, plan.anchor_pair, anchor_context |
| ) |
| ) |
| anchor_components = dict(anchor_context.get("_components") or {}) |
| block_score = float( |
| self.block_prior.score_block( |
| terminal.sequence_tokens, plan.anchor_pair, block, context |
| ) |
| ) |
| legal = math.isfinite(anchor_score) and math.isfinite(block_score) |
|
|
| logit = ( |
| self.cfg.geometry_coefficient * anchor_score |
| + self.cfg.esm2_coefficient * float(esm2_score) |
| ) / self.cfg.temperature |
| if not legal: |
| logit = float("-inf") |
| logits.append(float(logit)) |
| diagnostics.append( |
| { |
| "mode": f"{plan.ordered_pair}/i,i+{plan.spacing}", |
| "anchor_pair": list(plan.anchor_pair), |
| "block_id": plan.block_id, |
| "stapep_probability": self.mode_prior.probability( |
| (plan.ordered_pair, plan.spacing) |
| ), |
| "stapep_tilted_weight": self.mode_prior.tilted_weight( |
| (plan.ordered_pair, plan.spacing) |
| ), |
| "q_mode_target": q_mode[(plan.ordered_pair, plan.spacing)], |
| "anchor_score_raw": anchor_score, |
| "anchor_components": anchor_components, |
| "esm2_delta_raw": float(esm2_score), |
| "block_score_diagnostic_only": block_score, |
| "block_legal": bool(legal), |
| "within_mode_logit": float(logit), |
| } |
| ) |
|
|
| weights = [0.0] * len(plans) |
| for mode in feasible_modes: |
| indices = [ |
| index |
| for index, plan in enumerate(plans) |
| if (plan.ordered_pair, plan.spacing) == mode |
| ] |
| mode_logits = [logits[index] for index in indices] |
| finite = [math.isfinite(value) for value in mode_logits] |
| if not any(finite): |
| continue |
| masked = [value if ok else -1e30 for value, ok in zip(mode_logits, finite)] |
| within = self._softmax(masked) |
| for local_index, plan_index in enumerate(indices): |
| diagnostics[plan_index]["q_within_mode"] = float(within[local_index]) |
| weights[plan_index] = q_mode[mode] * within[local_index] |
| if len(indices) > 1: |
| weights[indices[-1]] = q_mode[mode] - sum( |
| weights[index] for index in indices[:-1] |
| ) |
|
|
| for index, weight in enumerate(weights): |
| diagnostics[index]["q_ref"] = float(weight) |
| self.last_diagnostics = diagnostics |
| return weights |
|
|