File size: 8,683 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
"""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