Download staplebridge/hydrocarbon/factorized_plan_reference.py from ChatterjeeLab/StapleBridge: direct link, hf CLI and curl.
- Browser
- Download file 8.68 kB
-
https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/hydrocarbon/factorized_plan_reference.py
- Command line
-
hf download hf://ChatterjeeLab/StapleBridge/staplebridge/hydrocarbon/factorized_plan_reference.py
-
curl -L -o factorized_plan_reference.py https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/hydrocarbon/factorized_plan_reference.py
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, | |
| ) | |
| 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 | |
| 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]] = [] | |
| 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 | |