"""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