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