from __future__ import annotations import math from typing import Any import torch from staplebridge.chemistry.state import StapleState from staplebridge.reference.energy import ( ACTION_ACTIVATE_TOPOLOGY, ACTION_ASSIGN_ANCHOR, ACTION_ASSIGN_BLOCK, ACTION_NOOP, ACTION_REASSIGN_ANCHOR, ACTION_RESIDUE_SUBSTITUTION_MOTIF, ACTION_RESIDUE_SUBSTITUTION_OTHER, ACTION_UNKNOWN, ReferenceEnergy, classify_action, ) from staplebridge.utils.profiling import STAGE_TIMER # Action groups for optional group-normalized sampling. Substitutions form one # group; anchor/block/topology actions each form their own so that a single # high-value transition is not swamped by a large fan-out of substitutions. ACTION_GROUP_MAP = { ACTION_NOOP: "noop", ACTION_RESIDUE_SUBSTITUTION_MOTIF: "substitution", ACTION_RESIDUE_SUBSTITUTION_OTHER: "substitution", ACTION_ASSIGN_ANCHOR: "anchor", ACTION_REASSIGN_ANCHOR: "anchor", ACTION_ASSIGN_BLOCK: "block", ACTION_ACTIVATE_TOPOLOGY: "topology", ACTION_UNKNOWN: "substitution", } class ReferenceKernel: def __init__( self, energy_model: ReferenceEnergy, group_normalize: bool = False, substitution_downweight: float = 1.0, ) -> None: self.energy_model = energy_model self.group_normalize = group_normalize # When both structural (anchor/block/topology) and substitution # candidates exist, multiply substitution probabilities by this factor # (0..1) before renormalizing. 1.0 = no downweight (default). self.substitution_downweight = float(substitution_downweight) def _decompose_all( self, z: StapleState, candidates: list[StapleState], context: dict[str, Any] | None, ) -> list[dict[str, Any]]: return self.energy_model.decompose_batch(z, candidates, context=context) def reference_logits( self, z: StapleState, candidates: list[StapleState], context: dict[str, Any] | None = None, ) -> torch.Tensor: with STAGE_TIMER.section("reference_logits_time"): decomps = self._decompose_all(z, candidates, context) out = torch.tensor([-d["E_total"] for d in decomps], dtype=torch.float32) STAGE_TIMER.bump("reference_logits_candidates", len(candidates)) return out def reference_probs( self, z: StapleState, candidates: list[StapleState], context: dict[str, Any] | None = None, ) -> torch.Tensor: with STAGE_TIMER.section("reference_logits_time"): decomps = self._decompose_all(z, candidates, context) logits = torch.tensor([-d["E_total"] for d in decomps], dtype=torch.float32) STAGE_TIMER.bump("reference_logits_candidates", len(candidates)) probs = torch.softmax(logits, dim=0) if self.group_normalize: # First sample action *group* uniformly over the groups actually # present in the neighborhood (weighted by aggregated group logit), # then softmax within group. Preserves the reference-energy # ordering while preventing large substitution fan-outs from # dominating the pmf. groups: dict[str, list[int]] = {} for idx, d in enumerate(decomps): g = ACTION_GROUP_MAP.get(d.get("action_type", ACTION_UNKNOWN), "substitution") groups.setdefault(g, []).append(idx) group_probs = probs.clone() group_probs.zero_() # Aggregate group score = logsumexp of member logits group_scores = {} for g, idxs in groups.items(): gl = logits[idxs] group_scores[g] = float(torch.logsumexp(gl, dim=0).item()) # Softmax over groups g_keys = list(group_scores.keys()) g_logits = torch.tensor([group_scores[k] for k in g_keys], dtype=torch.float32) g_pmf = torch.softmax(g_logits, dim=0) for gi, gk in enumerate(g_keys): idxs = groups[gk] sub_logits = logits[idxs] sub_pmf = torch.softmax(sub_logits, dim=0) group_probs[idxs] = sub_pmf * g_pmf[gi] probs = group_probs if self.substitution_downweight != 1.0: structural_present = any( ACTION_GROUP_MAP.get(d.get("action_type", ACTION_UNKNOWN), "substitution") in ("anchor", "block", "topology") for d in decomps ) if structural_present: factors = torch.tensor( [ self.substitution_downweight if ACTION_GROUP_MAP.get(d.get("action_type", ACTION_UNKNOWN), "substitution") == "substitution" else 1.0 for d in decomps ], dtype=torch.float32, ) probs = probs * factors s = probs.sum() if float(s.item()) > 0.0: probs = probs / s return probs def sample_next( self, z: StapleState, candidates: list[StapleState], context: dict[str, Any] | None = None, ) -> StapleState: probs = self.reference_probs(z, candidates, context=context) idx = torch.multinomial(probs, num_samples=1).item() return candidates[idx]