Download staplebridge/reference/kernel.py from ChatterjeeLab/StapleBridge: direct link, hf CLI and curl.
- Browser
- Download file 5.54 kB
-
https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/refs%2Fpr%2F1/staplebridge/reference/kernel.py
- Command line
-
hf download hf://ChatterjeeLab/StapleBridge@refs/pr/1/staplebridge/reference/kernel.py
-
curl -L -o kernel.py https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/refs%2Fpr%2F1/staplebridge/reference/kernel.py
5.54 kB
| 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] | |